1923 lines
81 KiB
Python
1923 lines
81 KiB
Python
"""Hierarchical, pageable memory components for Natural Memory v2.
|
|
|
|
The module deliberately keeps storage and neural routing separate:
|
|
|
|
* :class:`MemoryRouterV2` is the trainable sparse router.
|
|
* :class:`PagedMemoryBankV2` owns versioned pages and records.
|
|
* :class:`MemoryOSV2` coordinates writing, reading, quarantine and repair.
|
|
* :class:`KVBudgetManagerV2` models the hot-context budget.
|
|
|
|
The bank can run entirely in process memory for research, or be backed by a
|
|
checkpoint/page store by serializing ``export_payload``. No prompt text is
|
|
assembled by this module; it returns evidence records and routing traces to
|
|
the model adapter, which can inject only the selected evidence.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import math
|
|
import re
|
|
import time
|
|
from dataclasses import asdict, dataclass, field
|
|
from typing import TYPE_CHECKING, Any, Iterable, Optional, Sequence
|
|
|
|
import torch
|
|
from torch import Tensor, nn
|
|
import torch.nn.functional as F
|
|
|
|
if TYPE_CHECKING:
|
|
from .tiered_memory_store_v2 import TieredMemoryStoreV2
|
|
|
|
|
|
STATUS_ACTIVE = "active"
|
|
STATUS_SUPERSEDED = "superseded"
|
|
STATUS_RETRACTED = "retracted"
|
|
STATUS_QUARANTINED = "quarantined"
|
|
|
|
|
|
def _now() -> int:
|
|
return int(time.time())
|
|
|
|
|
|
def _stable_id(text: str, *, prefix: str) -> str:
|
|
digest = hashlib.sha1(text.encode("utf-8")).hexdigest()[:16]
|
|
return f"{prefix}_{digest}"
|
|
|
|
|
|
def _tokens(text: str) -> set[str]:
|
|
return {
|
|
token
|
|
for token in re.findall(r"[\u4e00-\u9fff]|[A-Za-z0-9_\-]+", text.lower())
|
|
if token not in {"我", "的", "是", "了", "请", "一下"}
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class MemoryRecordV2:
|
|
"""One versioned memory item with evidence and routing metadata."""
|
|
|
|
record_id: str
|
|
text: str
|
|
key: Tensor
|
|
summary: Tensor
|
|
memory_type: str = "fact"
|
|
entity: str = ""
|
|
attribute: str = ""
|
|
value: str = ""
|
|
timestamp: int = field(default_factory=_now)
|
|
importance: float = 0.5
|
|
confidence: float = 0.5
|
|
source: str = "user"
|
|
status: str = STATUS_ACTIVE
|
|
version: int = 0
|
|
page_id: str = ""
|
|
supersedes: str = ""
|
|
related_ids: list[str] = field(default_factory=list)
|
|
evidence: list[str] = field(default_factory=list)
|
|
# ``slot_index`` links a V2 address to the existing hot text bank when
|
|
# the Qwen adapter is used as a compatibility bridge. A value of -1
|
|
# means the record is an independent paged-memory item.
|
|
slot_index: int = -1
|
|
token_ids: Optional[Tensor] = None
|
|
token_mask: Optional[Tensor] = None
|
|
access_count: int = 0
|
|
last_access: int = field(default_factory=_now)
|
|
|
|
def conflict_key(self) -> str:
|
|
if self.entity and self.attribute:
|
|
return f"{self.entity.strip().lower()}::{self.attribute.strip().lower()}"
|
|
return ""
|
|
|
|
def lexical_score(self, query_text: str) -> float:
|
|
query = _tokens(query_text)
|
|
if not query:
|
|
return 0.0
|
|
own = _tokens(self.text)
|
|
if not own:
|
|
return 0.0
|
|
return len(query & own) / math.sqrt(max(1, len(query) * len(own)))
|
|
|
|
def token_overlap_score(self, query_token_ids: Optional[Tensor]) -> float:
|
|
"""Score exact token evidence without decoding or full attention."""
|
|
|
|
if query_token_ids is None or self.token_ids is None:
|
|
return 0.0
|
|
own_mask = self.token_mask
|
|
own = self.token_ids.reshape(-1)
|
|
if isinstance(own_mask, Tensor) and own_mask.numel() == own.numel():
|
|
own = own[own_mask.reshape(-1).to(dtype=torch.bool)]
|
|
query = query_token_ids.detach().reshape(-1).to(device=own.device)
|
|
if own.numel() == 0 or query.numel() == 0:
|
|
return 0.0
|
|
own_unique = torch.unique(own)
|
|
query_unique = torch.unique(query)
|
|
shared = torch.isin(query_unique, own_unique).sum().item()
|
|
# Query coverage is the useful direction here: a short question that
|
|
# names an entity should strongly prefer the chunk containing that
|
|
# entity even when the chunk itself is long.
|
|
return float(shared) / max(1, int(query_unique.numel()))
|
|
|
|
def touch(self) -> None:
|
|
self.access_count += 1
|
|
self.last_access = _now()
|
|
|
|
|
|
def memory_record_to_dict(record: MemoryRecordV2) -> dict[str, Any]:
|
|
"""Return a JSON-safe public view without exposing embedding tensors."""
|
|
|
|
return {
|
|
"record_id": record.record_id,
|
|
"text": record.text,
|
|
"memory_type": record.memory_type,
|
|
"entity": record.entity,
|
|
"attribute": record.attribute,
|
|
"value": record.value,
|
|
"timestamp": int(record.timestamp),
|
|
"importance": float(record.importance),
|
|
"confidence": float(record.confidence),
|
|
"source": record.source,
|
|
"status": record.status,
|
|
"version": int(record.version),
|
|
"page_id": record.page_id,
|
|
"supersedes": record.supersedes,
|
|
"related_ids": list(record.related_ids),
|
|
"evidence": list(record.evidence),
|
|
"slot_index": int(record.slot_index),
|
|
"token_count": int(record.token_ids.numel()) if isinstance(record.token_ids, Tensor) else 0,
|
|
"key_dim": int(record.key.numel()) if isinstance(record.key, Tensor) else 0,
|
|
"access_count": int(record.access_count),
|
|
"last_access": int(record.last_access),
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class MemoryPageV2:
|
|
page_id: str
|
|
tier: str = "warm"
|
|
capacity: int = 32
|
|
record_ids: list[str] = field(default_factory=list)
|
|
key: Optional[Tensor] = None
|
|
summary: Optional[Tensor] = None
|
|
importance: float = 0.0
|
|
created_at: int = field(default_factory=_now)
|
|
last_access: int = field(default_factory=_now)
|
|
|
|
@property
|
|
def full(self) -> bool:
|
|
return len(self.record_ids) >= self.capacity
|
|
|
|
def touch(self) -> None:
|
|
self.last_access = _now()
|
|
|
|
|
|
@dataclass
|
|
class RouterDecisionV2:
|
|
need_memory: bool
|
|
page_ids: list[str]
|
|
record_ids: list[str]
|
|
page_scores: list[float]
|
|
record_scores: list[float]
|
|
hop_count: int
|
|
hop_trace: list[list[str]]
|
|
confidence: float
|
|
stop_reason: str
|
|
|
|
|
|
class MemoryRouterV2(nn.Module):
|
|
"""Trainable multi-head router for sparse page and record retrieval.
|
|
|
|
The router does not generate the answer. It scores candidate memory
|
|
units, predicts whether another hop is useful, and exposes per-head scores
|
|
for diagnostics. The projection weights are trained separately from the
|
|
frozen Qwen backbone using hard-negative episodes.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
hidden_size: int,
|
|
*,
|
|
router_dim: int = 128,
|
|
num_heads: int = 8,
|
|
max_hops: int = 3,
|
|
) -> None:
|
|
super().__init__()
|
|
if router_dim % num_heads != 0:
|
|
raise ValueError("router_dim must be divisible by num_heads")
|
|
if max_hops < 1:
|
|
raise ValueError("max_hops must be positive")
|
|
self.hidden_size = hidden_size
|
|
self.router_dim = router_dim
|
|
self.num_heads = num_heads
|
|
self.max_hops = max_hops
|
|
self.head_dim = router_dim // num_heads
|
|
self.query_projection = nn.Linear(hidden_size, router_dim, bias=False)
|
|
self.key_projection = nn.Linear(hidden_size, router_dim, bias=False)
|
|
self.pair_scorer = nn.Sequential(
|
|
nn.Linear(router_dim * 3, router_dim),
|
|
nn.SiLU(),
|
|
nn.Linear(router_dim, 1),
|
|
)
|
|
policy_hidden = max(32, min(256, router_dim * 2))
|
|
self.need_memory = nn.Sequential(
|
|
nn.Linear(hidden_size, policy_hidden),
|
|
nn.SiLU(),
|
|
nn.Linear(policy_hidden, 1),
|
|
)
|
|
self.hop_controller = nn.Sequential(
|
|
nn.Linear(hidden_size, policy_hidden),
|
|
nn.SiLU(),
|
|
nn.Linear(policy_hidden, max_hops + 1),
|
|
)
|
|
self.head_gate = nn.Linear(hidden_size, num_heads)
|
|
|
|
def _reshape(self, value: Tensor) -> Tensor:
|
|
return value.view(*value.shape[:-1], self.num_heads, self.head_dim)
|
|
|
|
def encode_query(self, query: Tensor) -> Tensor:
|
|
"""Project a model hidden state into the compact address space."""
|
|
|
|
if query.shape[-1] != self.hidden_size:
|
|
raise ValueError(f"query last dimension must be {self.hidden_size}")
|
|
return F.normalize(self.query_projection(query), dim=-1)
|
|
|
|
def encode_key(self, key: Tensor) -> Tensor:
|
|
"""Project model keys into the compact address space."""
|
|
|
|
if key.shape[-1] != self.hidden_size:
|
|
raise ValueError(f"key last dimension must be {self.hidden_size}")
|
|
return F.normalize(self.key_projection(key), dim=-1)
|
|
|
|
def projected_scores(self, query: Tensor, projected_candidates: Tensor) -> tuple[Tensor, Tensor]:
|
|
"""Score compact candidates without storing full hidden states.
|
|
|
|
This is the storage-saving path: the bank keeps only the projected
|
|
candidate keys, while the current query is projected on demand.
|
|
"""
|
|
|
|
if query.ndim != 2 or query.shape[-1] != self.hidden_size:
|
|
raise ValueError("query must have shape [B, hidden_size]")
|
|
if projected_candidates.ndim == 2:
|
|
projected_candidates = projected_candidates.unsqueeze(0).expand(query.shape[0], -1, -1)
|
|
if projected_candidates.ndim != 3 or projected_candidates.shape[-1] != self.router_dim:
|
|
raise ValueError("projected_candidates must have shape [B,N,router_dim]")
|
|
q = self._reshape(self.encode_query(query))
|
|
k = self._reshape(F.normalize(projected_candidates, dim=-1))
|
|
head_scores = torch.einsum("bhc,bnhc->bnh", q, k)
|
|
gates = torch.softmax(self.head_gate(query), dim=-1)[:, None, :]
|
|
cosine_score = (head_scores * gates).sum(dim=-1)
|
|
q_expanded = q.reshape(query.shape[0], 1, -1).expand(-1, projected_candidates.shape[1], -1)
|
|
k_flat = k.reshape(query.shape[0], projected_candidates.shape[1], -1)
|
|
pair_input = torch.cat((q_expanded, k_flat, q_expanded - k_flat), dim=-1)
|
|
learned_score = self.pair_scorer(pair_input).squeeze(-1)
|
|
return cosine_score + learned_score, head_scores
|
|
|
|
def pair_scores(self, query: Tensor, candidates: Tensor) -> tuple[Tensor, Tensor]:
|
|
"""Return aggregate and per-head candidate scores.
|
|
|
|
``query`` is ``[B,H]`` and ``candidates`` is ``[B,N,H]`` or ``[N,H]``.
|
|
"""
|
|
|
|
if query.ndim != 2:
|
|
raise ValueError("query must have shape [B,H]")
|
|
if candidates.ndim == 2:
|
|
candidates = candidates.unsqueeze(0).expand(query.shape[0], -1, -1)
|
|
if candidates.ndim != 3 or candidates.shape[0] != query.shape[0]:
|
|
raise ValueError("candidates must have shape [B,N,H] or [N,H]")
|
|
return self.projected_scores(query, self.encode_key(candidates))
|
|
|
|
def forward(self, query: Tensor, candidates: Tensor) -> dict[str, Tensor]:
|
|
scores, head_scores = self.pair_scores(query, candidates)
|
|
return {
|
|
"scores": scores,
|
|
"head_scores": head_scores,
|
|
"need_memory_logits": self.need_memory(query).squeeze(-1),
|
|
"hop_logits": self.hop_controller(query),
|
|
}
|
|
|
|
|
|
class PagedMemoryBankV2:
|
|
"""Versioned memory pages with sparse routing and multi-hop expansion."""
|
|
|
|
def __init__(
|
|
self,
|
|
hidden_size: int,
|
|
*,
|
|
page_capacity: int = 32,
|
|
max_pages: int = 32768,
|
|
hot_pages: int = 8,
|
|
top_k_pages: int = 4,
|
|
top_k_records: int = 8,
|
|
max_hops: int = 3,
|
|
router: Optional[MemoryRouterV2] = None,
|
|
key_dim: Optional[int] = None,
|
|
coarse_index_bits: int = 20,
|
|
tier_store: Optional["TieredMemoryStoreV2"] = None,
|
|
max_resident_pages: int = 256,
|
|
runtime_device: Optional[torch.device] = None,
|
|
gpu_cache_records: int = 256,
|
|
gpu_cache_tokens: int = 131072,
|
|
gpu_cache_reserve_mb: int = 2048,
|
|
gpu_cache_adaptive: bool = True,
|
|
) -> None:
|
|
if page_capacity < 1 or max_pages < 1:
|
|
raise ValueError("page_capacity and max_pages must be positive")
|
|
self.hidden_size = hidden_size
|
|
self.page_capacity = page_capacity
|
|
self.max_pages = max_pages
|
|
self.hot_pages = hot_pages
|
|
self.top_k_pages = top_k_pages
|
|
self.top_k_records = top_k_records
|
|
self.max_hops = max_hops
|
|
self.router = router
|
|
self.key_dim = int(key_dim or (router.router_dim if router is not None else hidden_size))
|
|
if not 4 <= coarse_index_bits <= 20:
|
|
raise ValueError("coarse_index_bits must be between 4 and 20")
|
|
self.coarse_index_bits = int(coarse_index_bits)
|
|
if max_resident_pages < hot_pages:
|
|
raise ValueError("max_resident_pages must be >= hot_pages")
|
|
self.tier_store = tier_store
|
|
self.max_resident_pages = int(max_resident_pages)
|
|
if gpu_cache_records < 0 or gpu_cache_tokens < 0:
|
|
raise ValueError("gpu cache limits must be non-negative")
|
|
if gpu_cache_reserve_mb < 0:
|
|
raise ValueError("gpu_cache_reserve_mb must be non-negative")
|
|
self.gpu_cache_records = int(gpu_cache_records)
|
|
self.gpu_cache_tokens = int(gpu_cache_tokens)
|
|
self.gpu_cache_reserve_mb = int(gpu_cache_reserve_mb)
|
|
self.gpu_cache_reserve_bytes = self.gpu_cache_reserve_mb * 1024 * 1024
|
|
self.gpu_cache_adaptive = bool(gpu_cache_adaptive)
|
|
self._gpu_device: Optional[torch.device] = None
|
|
self._gpu_record_cache: dict[str, dict[str, Tensor]] = {}
|
|
self._gpu_cache_order: list[str] = []
|
|
self._gpu_cache_tokens_used = 0
|
|
self._gpu_cache_bytes_used = 0
|
|
self._gpu_cache_hits = 0
|
|
self._gpu_cache_misses = 0
|
|
self._gpu_cache_fallbacks = 0
|
|
self._gpu_cache_alloc_failures = 0
|
|
self._gpu_cache_last_free_bytes: Optional[int] = None
|
|
generator = torch.Generator(device="cpu").manual_seed(1729 + self.key_dim)
|
|
self.coarse_planes = F.normalize(
|
|
torch.randn(self.coarse_index_bits, self.key_dim, generator=generator), dim=-1
|
|
)
|
|
self.records: dict[str, MemoryRecordV2] = {}
|
|
self.pages: dict[str, MemoryPageV2] = {}
|
|
# Only pages with free capacity participate in write placement. The
|
|
# bounded order avoids scanning every page at million-record scale.
|
|
self._open_page_ids: set[str] = set()
|
|
self._open_page_order: list[str] = []
|
|
self._coarse_buckets: dict[int, set[str]] = {}
|
|
self._page_signatures: dict[str, set[int]] = {}
|
|
self._text_index: dict[str, str] = {}
|
|
self.active_by_conflict: dict[str, str] = {}
|
|
self.quarantine: dict[str, MemoryRecordV2] = {}
|
|
self._counter = 0
|
|
self._last_coarse_candidates = 0
|
|
self._suspend_refresh = 0
|
|
if self.tier_store is not None:
|
|
self._restore_page_headers_from_store()
|
|
for item in self.tier_store.load_quarantine():
|
|
record = MemoryRecordV2(**item)
|
|
self.quarantine[record.record_id] = record
|
|
self.active_by_conflict.update(
|
|
self.tier_store.active_conflicts(active_status=STATUS_ACTIVE)
|
|
)
|
|
if runtime_device is not None:
|
|
self.configure_gpu_cache(
|
|
runtime_device,
|
|
max_records=self.gpu_cache_records,
|
|
max_tokens=self.gpu_cache_tokens,
|
|
reserve_mb=self.gpu_cache_reserve_mb,
|
|
adaptive=self.gpu_cache_adaptive,
|
|
)
|
|
|
|
def configure_gpu_cache(
|
|
self,
|
|
device: torch.device | str,
|
|
*,
|
|
max_records: Optional[int] = None,
|
|
max_tokens: Optional[int] = None,
|
|
reserve_mb: Optional[int] = None,
|
|
adaptive: Optional[bool] = None,
|
|
) -> None:
|
|
"""Configure a bounded VRAM cache for recently used memory records.
|
|
|
|
The canonical copy remains in process RAM and is what gets exported
|
|
into the embedded third safetensors shard. This cache contains only
|
|
routing keys and the small token payloads needed by recent reads; it
|
|
is deliberately not a second full memory store.
|
|
"""
|
|
|
|
target = torch.device(device)
|
|
if max_records is not None:
|
|
if int(max_records) < 0:
|
|
raise ValueError("max_records must be non-negative")
|
|
self.gpu_cache_records = int(max_records)
|
|
if max_tokens is not None:
|
|
if int(max_tokens) < 0:
|
|
raise ValueError("max_tokens must be non-negative")
|
|
self.gpu_cache_tokens = int(max_tokens)
|
|
if reserve_mb is not None:
|
|
if int(reserve_mb) < 0:
|
|
raise ValueError("reserve_mb must be non-negative")
|
|
self.gpu_cache_reserve_mb = int(reserve_mb)
|
|
self.gpu_cache_reserve_bytes = self.gpu_cache_reserve_mb * 1024 * 1024
|
|
if adaptive is not None:
|
|
self.gpu_cache_adaptive = bool(adaptive)
|
|
self._gpu_device = target if target.type == "cuda" else None
|
|
self._clear_gpu_cache()
|
|
|
|
def _clear_gpu_cache(self) -> None:
|
|
self._gpu_record_cache.clear()
|
|
self._gpu_cache_order.clear()
|
|
self._gpu_cache_tokens_used = 0
|
|
self._gpu_cache_bytes_used = 0
|
|
|
|
@staticmethod
|
|
def _tensor_bytes(value: Optional[Tensor]) -> int:
|
|
if not isinstance(value, Tensor):
|
|
return 0
|
|
return int(value.numel()) * int(value.element_size())
|
|
|
|
def _record_cache_bytes(self, record: MemoryRecordV2, *, include_tokens: bool) -> int:
|
|
total = self._tensor_bytes(record.key) + self._tensor_bytes(record.summary)
|
|
if include_tokens:
|
|
total += self._tensor_bytes(record.token_ids)
|
|
total += self._tensor_bytes(record.token_mask)
|
|
return total
|
|
|
|
def _has_gpu_headroom(self, requested_bytes: int) -> bool:
|
|
"""Keep the cache below a VRAM safety line while the model is running."""
|
|
|
|
if self._gpu_device is None or not self.gpu_cache_adaptive:
|
|
return True
|
|
try:
|
|
free_bytes, _ = torch.cuda.mem_get_info(self._gpu_device)
|
|
except (RuntimeError, AssertionError, TypeError):
|
|
return True
|
|
self._gpu_cache_last_free_bytes = int(free_bytes)
|
|
return int(free_bytes) - int(requested_bytes) >= self.gpu_cache_reserve_bytes
|
|
|
|
def _evict_oldest_gpu_record(self) -> None:
|
|
if not self._gpu_cache_order:
|
|
return
|
|
evicted_id = self._gpu_cache_order.pop(0)
|
|
evicted = self._gpu_record_cache.pop(evicted_id, {})
|
|
evicted_tokens = evicted.get("token_ids")
|
|
if isinstance(evicted_tokens, Tensor):
|
|
self._gpu_cache_tokens_used -= int(evicted_tokens.numel())
|
|
self._gpu_cache_bytes_used -= sum(self._tensor_bytes(value) for value in evicted.values())
|
|
self._gpu_cache_bytes_used = max(0, self._gpu_cache_bytes_used)
|
|
|
|
def _promote_record(self, record: MemoryRecordV2) -> None:
|
|
"""Promote one hot record to VRAM without changing its RAM copy."""
|
|
|
|
if (
|
|
self._gpu_device is None
|
|
or self.gpu_cache_records <= 0
|
|
or record.record_id in self._gpu_record_cache
|
|
):
|
|
if record.record_id in self._gpu_record_cache:
|
|
self._gpu_cache_hits += 1
|
|
self._gpu_cache_order.remove(record.record_id)
|
|
self._gpu_cache_order.append(record.record_id)
|
|
return
|
|
self._gpu_cache_misses += 1
|
|
source_token_count = int(record.token_ids.numel()) if isinstance(record.token_ids, Tensor) else 0
|
|
token_count = source_token_count
|
|
if self.gpu_cache_tokens == 0 or token_count > self.gpu_cache_tokens:
|
|
token_count = 0
|
|
include_tokens = token_count > 0 and isinstance(record.token_ids, Tensor)
|
|
requested_bytes = self._record_cache_bytes(record, include_tokens=include_tokens)
|
|
while self._gpu_cache_order and (
|
|
len(self._gpu_cache_order) >= self.gpu_cache_records
|
|
or self._gpu_cache_tokens_used + token_count > self.gpu_cache_tokens
|
|
or not self._has_gpu_headroom(requested_bytes)
|
|
):
|
|
self._evict_oldest_gpu_record()
|
|
if not self._has_gpu_headroom(requested_bytes):
|
|
# The authoritative RAM copy remains available to the router.
|
|
# A hot-cache miss must never make inference fail just because the
|
|
# GPU is busy with model weights or KV activations.
|
|
self._gpu_cache_fallbacks += 1
|
|
return
|
|
try:
|
|
cached: dict[str, Tensor] = {
|
|
"key": record.key.detach().to(self._gpu_device, non_blocking=True),
|
|
"summary": record.summary.detach().to(self._gpu_device, non_blocking=True),
|
|
}
|
|
if include_tokens:
|
|
cached["token_ids"] = record.token_ids.detach().to(self._gpu_device, non_blocking=True)
|
|
if isinstance(record.token_mask, Tensor):
|
|
cached["token_mask"] = record.token_mask.detach().to(self._gpu_device, non_blocking=True)
|
|
except RuntimeError as exc:
|
|
if "out of memory" not in str(exc).lower():
|
|
raise
|
|
self._gpu_cache_alloc_failures += 1
|
|
self._gpu_cache_fallbacks += 1
|
|
return
|
|
if include_tokens:
|
|
self._gpu_cache_tokens_used += token_count
|
|
cache_bytes = sum(self._tensor_bytes(value) for value in cached.values())
|
|
self._gpu_record_cache[record.record_id] = cached
|
|
self._gpu_cache_order.append(record.record_id)
|
|
self._gpu_cache_bytes_used += cache_bytes
|
|
|
|
def _record_key(self, record: MemoryRecordV2) -> Tensor:
|
|
self._promote_record(record)
|
|
cached = self._gpu_record_cache.get(record.record_id)
|
|
if cached is not None:
|
|
return cached["key"]
|
|
return record.key
|
|
|
|
def gpu_record_payload(self, record: MemoryRecordV2) -> tuple[Tensor, Optional[Tensor]]:
|
|
"""Return a hot record's token payload, preferring its VRAM copy."""
|
|
|
|
self._promote_record(record)
|
|
cached = self._gpu_record_cache.get(record.record_id)
|
|
if cached is not None and "token_ids" in cached:
|
|
return cached["token_ids"], cached.get("token_mask")
|
|
return record.token_ids, record.token_mask
|
|
|
|
def promote_records(self, records: Iterable[MemoryRecordV2]) -> None:
|
|
for record in records:
|
|
self._promote_record(record)
|
|
|
|
def has_token_evidence(
|
|
self,
|
|
query_key: Tensor,
|
|
query_token_ids: Optional[Tensor],
|
|
*,
|
|
min_shared_tokens: int = 3,
|
|
) -> bool:
|
|
"""Check bounded exact evidence before the learned abstention gate."""
|
|
|
|
if query_token_ids is None:
|
|
return False
|
|
query_unique = torch.unique(query_token_ids.detach().reshape(-1).cpu())
|
|
if query_unique.numel() < min_shared_tokens:
|
|
return False
|
|
for page_id in self._candidate_page_ids(query_key):
|
|
page = self._hydrate_page(page_id)
|
|
if page is None:
|
|
continue
|
|
for record_id in page.record_ids:
|
|
record = self.records.get(record_id)
|
|
if record is None or record.status != STATUS_ACTIVE or record.token_ids is None:
|
|
continue
|
|
own = record.token_ids.reshape(-1)
|
|
if isinstance(record.token_mask, Tensor) and record.token_mask.numel() == own.numel():
|
|
own = own[record.token_mask.reshape(-1).to(dtype=torch.bool)]
|
|
shared = torch.isin(query_unique, torch.unique(own)).sum().item()
|
|
if int(shared) >= int(min_shared_tokens):
|
|
return True
|
|
return False
|
|
|
|
def _restore_page_headers_from_store(self) -> None:
|
|
"""Restore page addresses without materializing every cold record."""
|
|
|
|
if self.tier_store is None:
|
|
return
|
|
headers = self.tier_store.page_headers()
|
|
for item in headers:
|
|
page = MemoryPageV2(**item)
|
|
self.pages[page.page_id] = page
|
|
self._counter = max(self._counter, int(page.page_id.rsplit("_", 1)[-1]))
|
|
if not page.full:
|
|
self._open_page_ids.add(page.page_id)
|
|
self._open_page_order.append(page.page_id)
|
|
self._index_page(page, persist=False)
|
|
|
|
def _rebuild_coarse_index(self, coarse_index_bits: Optional[int] = None) -> None:
|
|
"""Recreate LSH planes and page buckets after a config upgrade."""
|
|
|
|
if coarse_index_bits is not None:
|
|
if not 4 <= int(coarse_index_bits) <= 20:
|
|
raise ValueError("coarse_index_bits must be between 4 and 20")
|
|
self.coarse_index_bits = int(coarse_index_bits)
|
|
generator = torch.Generator(device="cpu").manual_seed(1729 + self.key_dim)
|
|
self.coarse_planes = F.normalize(
|
|
torch.randn(self.coarse_index_bits, self.key_dim, generator=generator), dim=-1
|
|
)
|
|
self._coarse_buckets.clear()
|
|
self._page_signatures.clear()
|
|
for page in self.pages.values():
|
|
self._index_page(page, persist=False)
|
|
if self.tier_store is not None:
|
|
self.tier_store.clear_coarse_buckets()
|
|
page_by_record = {
|
|
record_id: page.page_id
|
|
for page in self.pages.values()
|
|
for record_id in page.record_ids
|
|
}
|
|
for record_id, key in self.tier_store.record_keys():
|
|
page_id = page_by_record.get(record_id)
|
|
if page_id is None:
|
|
continue
|
|
signature = self._coarse_signature(key)
|
|
self._page_signatures.setdefault(page_id, set()).add(signature)
|
|
for page_id, signatures in self._page_signatures.items():
|
|
self.tier_store.replace_page_buckets(page_id, signatures)
|
|
|
|
def _encode_key(self, key: Tensor) -> Tensor:
|
|
key = key.detach().float().reshape(-1)
|
|
if key.numel() == self.key_dim:
|
|
return F.normalize(key, dim=0)
|
|
if self.router is not None and key.numel() == self.hidden_size:
|
|
router_device = next(self.router.parameters()).device
|
|
return self.router.encode_key(key.to(router_device).unsqueeze(0))[0].detach().float().cpu()
|
|
raise ValueError(
|
|
f"memory key must contain {self.hidden_size} model values or {self.key_dim} address values"
|
|
)
|
|
|
|
def _score_candidates(self, query_key: Tensor, candidate_keys: Tensor) -> Tensor:
|
|
if self.router is not None and query_key.numel() == self.hidden_size:
|
|
router_device = next(self.router.parameters()).device
|
|
scores, _ = self.router.projected_scores(
|
|
query_key.to(router_device).reshape(1, -1),
|
|
candidate_keys.to(router_device).reshape(1, -1, self.key_dim),
|
|
)
|
|
# Training uses logits for cross-entropy; storage uses a bounded
|
|
# relevance value so page thresholds remain stable across router
|
|
# checkpoints and do not discard valid negative logits.
|
|
return torch.sigmoid(scores[0])
|
|
query = self._encode_key(query_key)
|
|
candidates = F.normalize(candidate_keys.float().to(query.device), dim=-1)
|
|
return torch.matmul(candidates, query)
|
|
|
|
def _coarse_signature(self, key: Tensor) -> int:
|
|
address = self._encode_key(key).cpu()
|
|
bits = (torch.mv(self.coarse_planes, address) >= 0).to(torch.int64)
|
|
signature = 0
|
|
for index, bit in enumerate(bits.tolist()):
|
|
signature |= int(bit) << index
|
|
return signature
|
|
|
|
def _remove_page_from_index(self, page: MemoryPageV2) -> None:
|
|
signatures = self._page_signatures.pop(page.page_id, set())
|
|
for signature in signatures:
|
|
bucket = self._coarse_buckets.get(signature)
|
|
if bucket is not None:
|
|
bucket.discard(page.page_id)
|
|
if not bucket:
|
|
self._coarse_buckets.pop(signature, None)
|
|
|
|
def _index_page(self, page: MemoryPageV2, *, persist: bool = True) -> None:
|
|
signatures: set[int] = set()
|
|
if page.key is not None:
|
|
signatures.add(self._coarse_signature(page.key))
|
|
# Index record addresses as well as the page centroid. A centroid
|
|
# can blur multiple topics, while a record address is exact and still
|
|
# costs only one small LSH bucket insertion per record.
|
|
for record_id in page.record_ids:
|
|
record = self.records.get(record_id)
|
|
if record is not None:
|
|
signatures.add(self._coarse_signature(record.key))
|
|
self._page_signatures[page.page_id] = signatures
|
|
for signature in signatures:
|
|
self._coarse_buckets.setdefault(signature, set()).add(page.page_id)
|
|
if self.tier_store is not None and persist:
|
|
self.tier_store.upsert_page(page)
|
|
self.tier_store.replace_page_buckets(page.page_id, signatures)
|
|
|
|
def _hydrate_records(self, record_ids: Iterable[str]) -> None:
|
|
"""Load only the records needed by the current candidate pages."""
|
|
|
|
if self.tier_store is None:
|
|
return
|
|
missing = [record_id for record_id in record_ids if record_id not in self.records]
|
|
if not missing:
|
|
return
|
|
for item in self.tier_store.load_records(missing):
|
|
self.records[item["record_id"]] = MemoryRecordV2(**item)
|
|
text = item.get("text", "").strip().lower()
|
|
if text and item.get("status") == STATUS_ACTIVE:
|
|
self._text_index[text] = item["record_id"]
|
|
|
|
def _hydrate_page(self, page_id: str) -> Optional[MemoryPageV2]:
|
|
page = self.pages.get(page_id)
|
|
if page is None or self.tier_store is None:
|
|
return page
|
|
self._hydrate_records(page.record_ids)
|
|
return page
|
|
|
|
def _evict_cold_records(self) -> None:
|
|
"""Keep page headers resident and unload non-hot record payloads."""
|
|
|
|
if self.tier_store is None or self.max_resident_pages <= 0:
|
|
return
|
|
ranked = sorted(
|
|
self.pages.values(),
|
|
key=lambda page: (page.tier == "hot", page.last_access, page.importance),
|
|
reverse=True,
|
|
)
|
|
keep = {page.page_id for page in ranked[: self.max_resident_pages]}
|
|
for page in self.pages.values():
|
|
if page.page_id in keep:
|
|
continue
|
|
for record_id in page.record_ids:
|
|
self.records.pop(record_id, None)
|
|
|
|
def _store_record(self, record: MemoryRecordV2) -> None:
|
|
if self.tier_store is not None:
|
|
self.tier_store.upsert_record(record)
|
|
|
|
def _candidate_page_ids(self, query_key: Tensor) -> list[str]:
|
|
"""Use locality-sensitive coarse buckets before exact page scoring."""
|
|
|
|
page_ids = list(self.pages)
|
|
if len(page_ids) <= 128:
|
|
self._last_coarse_candidates = len(page_ids)
|
|
return page_ids
|
|
signature = self._coarse_signature(query_key)
|
|
probes = [signature]
|
|
# Probe Hamming distance 1 and 2. This is bounded (<= 79 buckets for
|
|
# 12 bits) and avoids a query-time scan over every page.
|
|
for bit in range(self.coarse_index_bits):
|
|
probes.append(signature ^ (1 << bit))
|
|
for left in range(self.coarse_index_bits):
|
|
for right in range(left + 1, self.coarse_index_bits):
|
|
probes.append(signature ^ (1 << left) ^ (1 << right))
|
|
if self.tier_store is not None:
|
|
hot_ids = [
|
|
page.page_id for page in self.pages.values() if page.tier == "hot"
|
|
]
|
|
selected = self.tier_store.candidate_page_ids(
|
|
probes,
|
|
hot_page_ids=hot_ids,
|
|
limit=max(128, self.top_k_pages * 64),
|
|
)
|
|
self._last_coarse_candidates = len(selected)
|
|
return selected
|
|
selected: set[str] = set()
|
|
for bucket in probes:
|
|
selected.update(self._coarse_buckets.get(bucket, ()))
|
|
# Hot pages are always eligible. They are few by construction and
|
|
# protect high-value memories when a page centroid is still moving.
|
|
selected.update(
|
|
page.page_id for page in self.pages.values() if page.tier == "hot"
|
|
)
|
|
if not selected:
|
|
# Cold-start safety: a bounded sample is preferable to silently
|
|
# claiming that no memory exists.
|
|
selected.update(page_ids[: min(128, len(page_ids))])
|
|
self._last_coarse_candidates = len(selected)
|
|
return list(selected)
|
|
|
|
def _new_page(self, *, tier: str = "warm") -> MemoryPageV2:
|
|
if len(self.pages) >= self.max_pages:
|
|
self.consolidate(max_pages=max(1, self.max_pages - 1))
|
|
if len(self.pages) >= self.max_pages:
|
|
raise RuntimeError(
|
|
"memory page capacity exhausted; consolidate or increase max_pages"
|
|
)
|
|
self._counter += 1
|
|
page = MemoryPageV2(
|
|
page_id=f"page_{self._counter:08d}",
|
|
tier=tier,
|
|
capacity=self.page_capacity,
|
|
)
|
|
self.pages[page.page_id] = page
|
|
self._open_page_ids.add(page.page_id)
|
|
self._open_page_order.append(page.page_id)
|
|
return page
|
|
|
|
def _target_page(self, key: Tensor, *, importance: float) -> MemoryPageV2:
|
|
open_ids = [
|
|
page_id
|
|
for page_id in self._open_page_order
|
|
if page_id in self._open_page_ids
|
|
and page_id in self.pages
|
|
and not self.pages[page_id].full
|
|
]
|
|
# If many pages remain partially filled, compare only a bounded recent
|
|
# window plus hot pages. Write placement must stay sublinear in the
|
|
# number of pages.
|
|
if len(open_ids) > 128:
|
|
recent = open_ids[-128:]
|
|
hot = [
|
|
page.page_id
|
|
for page in self.pages.values()
|
|
if page.tier == "hot" and not page.full
|
|
]
|
|
open_ids = list(dict.fromkeys(recent + hot))
|
|
candidates = [self.pages[page_id] for page_id in open_ids]
|
|
if not candidates:
|
|
return self._new_page(tier="hot" if importance >= 0.8 else "warm")
|
|
key = self._encode_key(key)
|
|
best_page: Optional[MemoryPageV2] = None
|
|
best_score = -float("inf")
|
|
for page in candidates:
|
|
if page.key is None:
|
|
score = -0.1 * len(page.record_ids)
|
|
else:
|
|
score = float(torch.dot(key.to(page.key.device), F.normalize(page.key.float(), dim=0)))
|
|
score += 0.02 * page.importance
|
|
if score > best_score:
|
|
best_score = score
|
|
best_page = page
|
|
assert best_page is not None
|
|
return best_page
|
|
|
|
def _update_page(self, page: MemoryPageV2, record: MemoryRecordV2) -> None:
|
|
self._remove_page_from_index(page)
|
|
page.record_ids.append(record.record_id)
|
|
record.page_id = page.page_id
|
|
weight = 1.0 / max(1, len(page.record_ids))
|
|
if page.key is None:
|
|
page.key = record.key.detach().float().clone()
|
|
page.summary = record.summary.detach().float().clone()
|
|
else:
|
|
page.key = (1.0 - weight) * page.key + weight * record.key.detach().float()
|
|
if page.summary is None:
|
|
page.summary = record.summary.detach().float().clone()
|
|
else:
|
|
page.summary = (1.0 - weight) * page.summary + weight * record.summary.detach().float()
|
|
page.importance = max(page.importance, float(record.importance))
|
|
page.touch()
|
|
if page.full:
|
|
self._open_page_ids.discard(page.page_id)
|
|
self._index_page(page)
|
|
|
|
def _find_duplicate(self, record: MemoryRecordV2) -> Optional[MemoryRecordV2]:
|
|
text_key = record.text.strip().lower()
|
|
if text_key:
|
|
record_id = self._text_index.get(text_key)
|
|
duplicate = self.records.get(record_id) if record_id else None
|
|
if duplicate is None and self.tier_store is not None:
|
|
item = self.tier_store.find_by_text(record.text, active_status=STATUS_ACTIVE)
|
|
if item is not None:
|
|
self._hydrate_records([item["record_id"]])
|
|
duplicate = self.records.get(item["record_id"])
|
|
if duplicate is not None and duplicate.status == STATUS_ACTIVE:
|
|
return duplicate
|
|
conflict_key = record.conflict_key()
|
|
if conflict_key:
|
|
old_id = self.active_by_conflict.get(conflict_key)
|
|
if old_id and old_id in self.records:
|
|
old = self.records[old_id]
|
|
if old.status == STATUS_ACTIVE and old.value.strip().lower() == record.value.strip().lower():
|
|
return old
|
|
if self.tier_store is not None:
|
|
item = self.tier_store.find_by_conflict(
|
|
record.entity,
|
|
record.attribute,
|
|
active_status=STATUS_ACTIVE,
|
|
)
|
|
if item is not None:
|
|
self._hydrate_records([item["record_id"]])
|
|
old = self.records.get(item["record_id"])
|
|
if old is not None and old.value.strip().lower() == record.value.strip().lower():
|
|
return old
|
|
return None
|
|
|
|
def write(
|
|
self,
|
|
*,
|
|
text: str,
|
|
key: Tensor,
|
|
summary: Optional[Tensor] = None,
|
|
memory_type: str = "fact",
|
|
entity: str = "",
|
|
attribute: str = "",
|
|
value: str = "",
|
|
importance: float = 0.5,
|
|
confidence: float = 0.5,
|
|
source: str = "user",
|
|
evidence: Optional[Sequence[str]] = None,
|
|
related_ids: Optional[Sequence[str]] = None,
|
|
slot_index: int = -1,
|
|
token_ids: Optional[Tensor] = None,
|
|
token_mask: Optional[Tensor] = None,
|
|
trusted: bool = True,
|
|
force: bool = False,
|
|
) -> tuple[MemoryRecordV2, str]:
|
|
"""Insert or version a record, returning ``(record, action)``."""
|
|
|
|
importance = float(max(0.0, min(1.0, importance)))
|
|
confidence = float(max(0.0, min(1.0, confidence)))
|
|
if summary is None:
|
|
summary = key
|
|
record = MemoryRecordV2(
|
|
record_id=_stable_id(f"{text}|{_now()}|{len(self.records)}", prefix="mem"),
|
|
text=text,
|
|
key=self._encode_key(key).cpu().clone(),
|
|
summary=self._encode_key(summary).cpu().clone(),
|
|
memory_type=memory_type,
|
|
entity=entity,
|
|
attribute=attribute,
|
|
value=value,
|
|
importance=importance,
|
|
confidence=confidence,
|
|
source=source,
|
|
evidence=list(evidence or [text]),
|
|
related_ids=list(related_ids or []),
|
|
slot_index=int(slot_index),
|
|
token_ids=(token_ids.detach().cpu().long().reshape(-1).clone() if token_ids is not None else None),
|
|
token_mask=(token_mask.detach().cpu().bool().reshape(-1).clone() if token_mask is not None else None),
|
|
)
|
|
if not trusted and not force:
|
|
record.status = STATUS_QUARANTINED
|
|
self.quarantine[record.record_id] = record
|
|
if self.tier_store is not None:
|
|
self.tier_store.upsert_quarantine(record)
|
|
return record, "quarantined"
|
|
duplicate = self._find_duplicate(record)
|
|
if duplicate is not None:
|
|
duplicate.confidence = max(duplicate.confidence, record.confidence)
|
|
duplicate.importance = max(duplicate.importance, record.importance)
|
|
duplicate.evidence.extend(item for item in record.evidence if item not in duplicate.evidence)
|
|
duplicate.touch()
|
|
return duplicate, "duplicate"
|
|
if record.slot_index >= 0:
|
|
for old in self.records.values():
|
|
if old.slot_index == record.slot_index and old.status == STATUS_ACTIVE:
|
|
old.status = STATUS_SUPERSEDED
|
|
record.version = max(record.version, old.version + 1)
|
|
record.supersedes = old.record_id
|
|
self._store_record(old)
|
|
conflict_key = record.conflict_key()
|
|
if conflict_key and conflict_key in self.active_by_conflict:
|
|
old_id = self.active_by_conflict[conflict_key]
|
|
old = self.records.get(old_id)
|
|
if old is not None and old.status == STATUS_ACTIVE:
|
|
old.status = STATUS_SUPERSEDED
|
|
record.version = old.version + 1
|
|
record.supersedes = old.record_id
|
|
self._store_record(old)
|
|
elif conflict_key and self.tier_store is not None:
|
|
item = self.tier_store.find_by_conflict(
|
|
record.entity,
|
|
record.attribute,
|
|
active_status=STATUS_ACTIVE,
|
|
)
|
|
old = None
|
|
if item is not None:
|
|
self._hydrate_records([item["record_id"]])
|
|
old = self.records.get(item["record_id"])
|
|
if old is not None:
|
|
old.status = STATUS_SUPERSEDED
|
|
record.version = old.version + 1
|
|
record.supersedes = old.record_id
|
|
self.active_by_conflict[conflict_key] = old.record_id
|
|
self._store_record(old)
|
|
page = self._target_page(record.key, importance=record.importance)
|
|
self.records[record.record_id] = record
|
|
if record.text.strip():
|
|
self._text_index[record.text.strip().lower()] = record.record_id
|
|
self._update_page(page, record)
|
|
self._store_record(record)
|
|
if conflict_key:
|
|
self.active_by_conflict[conflict_key] = record.record_id
|
|
if self._suspend_refresh == 0:
|
|
self._refresh_tiers()
|
|
return record, "updated" if record.supersedes else "inserted"
|
|
|
|
def write_batch(self, records: Iterable[dict[str, Any]]) -> list[tuple[MemoryRecordV2, str]]:
|
|
"""Write many records with one durable transaction and one tier pass."""
|
|
|
|
items = list(records)
|
|
if not items:
|
|
return []
|
|
self._suspend_refresh += 1
|
|
try:
|
|
if self.tier_store is not None:
|
|
with self.tier_store.transaction():
|
|
output = [self.write(**item) for item in items]
|
|
else:
|
|
output = [self.write(**item) for item in items]
|
|
finally:
|
|
self._suspend_refresh -= 1
|
|
if self._suspend_refresh == 0:
|
|
self._refresh_tiers()
|
|
return output
|
|
|
|
def _refresh_tiers(self) -> None:
|
|
ranked = sorted(
|
|
self.pages.values(),
|
|
key=lambda page: (page.importance, page.last_access, len(page.record_ids)),
|
|
reverse=True,
|
|
)
|
|
hot_ids = {page.page_id for page in ranked[: self.hot_pages]}
|
|
resident_ids = {
|
|
page.page_id for page in ranked[: self.max_resident_pages]
|
|
}
|
|
for page in self.pages.values():
|
|
old_tier = page.tier
|
|
if page.page_id in hot_ids:
|
|
page.tier = "hot"
|
|
elif self.tier_store is not None and page.page_id not in resident_ids:
|
|
page.tier = "cold"
|
|
else:
|
|
page.tier = "warm"
|
|
if self.tier_store is not None and old_tier != page.tier:
|
|
self.tier_store.set_page_tier(page.page_id, page.tier)
|
|
if self.tier_store is None and self._gpu_device is not None:
|
|
hot_records: list[MemoryRecordV2] = []
|
|
for page in sorted(
|
|
(item for item in self.pages.values() if item.tier == "hot"),
|
|
key=lambda item: (item.last_access, item.importance),
|
|
reverse=True,
|
|
):
|
|
hot_records.extend(
|
|
self.records[record_id]
|
|
for record_id in page.record_ids
|
|
if record_id in self.records
|
|
and self.records[record_id].status == STATUS_ACTIVE
|
|
)
|
|
self.promote_records(hot_records)
|
|
self._evict_cold_records()
|
|
|
|
def _page_scores(
|
|
self,
|
|
query_key: Tensor,
|
|
query_text: str = "",
|
|
query_token_ids: Optional[Tensor] = None,
|
|
) -> list[tuple[MemoryPageV2, float]]:
|
|
scores = []
|
|
candidate_page_ids = self._candidate_page_ids(query_key)
|
|
for page_id in candidate_page_ids:
|
|
page = self._hydrate_page(page_id)
|
|
if page is None:
|
|
continue
|
|
if page.key is None:
|
|
continue
|
|
score = float(self._score_candidates(query_key, page.key.unsqueeze(0))[0].item())
|
|
active_records = [
|
|
self.records[record_id]
|
|
for record_id in page.record_ids
|
|
if record_id in self.records and self.records[record_id].status == STATUS_ACTIVE
|
|
]
|
|
if active_records:
|
|
record_keys = torch.stack([self._record_key(record) for record in active_records], dim=0)
|
|
# Exact record-level reranking prevents a mixed-topic page
|
|
# centroid from hiding a highly relevant slot.
|
|
score = max(score, float(self._score_candidates(query_key, record_keys).max().item()))
|
|
score += 0.05 * page.importance
|
|
age_seconds = max(0, _now() - page.last_access)
|
|
score += 0.02 / (1.0 + age_seconds / 3600.0)
|
|
if query_text:
|
|
page_text = " ".join(self.records[item].text for item in page.record_ids if item in self.records)
|
|
score += 0.25 * max(
|
|
(self.records[item].lexical_score(query_text) for item in page.record_ids if item in self.records),
|
|
default=0.0,
|
|
)
|
|
if query_token_ids is not None:
|
|
score += 0.45 * max(
|
|
(
|
|
self.records[item].token_overlap_score(query_token_ids)
|
|
for item in page.record_ids
|
|
if item in self.records and self.records[item].status == STATUS_ACTIVE
|
|
),
|
|
default=0.0,
|
|
)
|
|
scores.append((page, score))
|
|
return sorted(scores, key=lambda item: item[1], reverse=True)
|
|
|
|
def _record_scores(
|
|
self,
|
|
page: MemoryPageV2,
|
|
query_key: Tensor,
|
|
query_text: str,
|
|
query_token_ids: Optional[Tensor] = None,
|
|
*,
|
|
allow_superseded: bool = False,
|
|
) -> list[tuple[MemoryRecordV2, float]]:
|
|
output = []
|
|
self._hydrate_page(page.page_id)
|
|
live_records: list[MemoryRecordV2] = []
|
|
for record_id in page.record_ids:
|
|
record = self.records.get(record_id)
|
|
if record is None:
|
|
continue
|
|
if record.status != STATUS_ACTIVE and not allow_superseded:
|
|
continue
|
|
live_records.append(record)
|
|
if not live_records:
|
|
return []
|
|
candidate_keys = torch.stack([self._record_key(record) for record in live_records], dim=0)
|
|
routed_scores = self._score_candidates(query_key, candidate_keys)
|
|
for record, routed_score in zip(live_records, routed_scores):
|
|
score = float(routed_score.item())
|
|
score += 0.25 * record.lexical_score(query_text)
|
|
score += 0.45 * record.token_overlap_score(query_token_ids)
|
|
score += 0.10 * record.confidence + 0.08 * record.importance
|
|
output.append((record, score))
|
|
return sorted(output, key=lambda item: item[1], reverse=True)
|
|
|
|
def query(
|
|
self,
|
|
*,
|
|
query_key: Tensor,
|
|
query_text: str = "",
|
|
query_token_ids: Optional[Tensor] = None,
|
|
top_k_pages: Optional[int] = None,
|
|
top_k_records: Optional[int] = None,
|
|
max_hops: Optional[int] = None,
|
|
min_score: float = -1.0,
|
|
) -> tuple[list[MemoryRecordV2], RouterDecisionV2]:
|
|
"""Sparse multi-hop query over pages and active records."""
|
|
|
|
page_limit = max(1, top_k_pages or self.top_k_pages)
|
|
record_limit = top_k_records or self.top_k_records
|
|
hop_limit = min(max_hops or self.max_hops, self.max_hops)
|
|
ranked_pages = self._page_scores(query_key, query_text, query_token_ids)
|
|
selected_pages = [item for item in ranked_pages[:page_limit] if item[1] >= min_score]
|
|
current_page_ids = [page.page_id for page, _ in selected_pages]
|
|
selected_records: list[tuple[MemoryRecordV2, float]] = []
|
|
hop_trace: list[list[str]] = []
|
|
visited: set[str] = set()
|
|
for hop in range(hop_limit):
|
|
hop_records: list[str] = []
|
|
candidate_pages = [self.pages[item] for item in current_page_ids if item in self.pages]
|
|
candidates = []
|
|
for page in candidate_pages:
|
|
candidates.extend(self._record_scores(page, query_key, query_text, query_token_ids))
|
|
candidates.sort(key=lambda item: item[1], reverse=True)
|
|
for record, score in candidates:
|
|
if record.record_id in visited:
|
|
continue
|
|
visited.add(record.record_id)
|
|
selected_records.append((record, score))
|
|
record.touch()
|
|
self._store_record(record)
|
|
hop_records.append(record.record_id)
|
|
if len(selected_records) >= record_limit:
|
|
break
|
|
hop_trace.append(hop_records)
|
|
if len(selected_records) >= record_limit or not hop_records:
|
|
break
|
|
related_pages: list[str] = []
|
|
for record_id in hop_records:
|
|
record = self.records[record_id]
|
|
self._hydrate_records(record.related_ids)
|
|
for related_id in record.related_ids:
|
|
related = self.records.get(related_id)
|
|
if related is not None and related.page_id not in related_pages:
|
|
related_pages.append(related.page_id)
|
|
if not related_pages:
|
|
break
|
|
current_page_ids = related_pages[:page_limit]
|
|
selected_records.sort(key=lambda item: item[1], reverse=True)
|
|
records = [record for record, _ in selected_records[:record_limit]]
|
|
# Keep only the bounded hot working set on VRAM. The authoritative
|
|
# copies remain in RAM and are exported to the embedded weight shard.
|
|
self.promote_records(records)
|
|
page_scores = [score for _, score in selected_pages]
|
|
record_scores = [score for _, score in selected_records[:record_limit]]
|
|
confidence = float(max(record_scores, default=0.0))
|
|
confidence = max(0.0, min(1.0, (confidence + 1.0) / 2.0))
|
|
decision = RouterDecisionV2(
|
|
need_memory=bool(records),
|
|
page_ids=[page.page_id for page, _ in selected_pages],
|
|
record_ids=[record.record_id for record in records],
|
|
page_scores=page_scores,
|
|
record_scores=record_scores,
|
|
hop_count=len(hop_trace),
|
|
hop_trace=hop_trace,
|
|
confidence=confidence,
|
|
stop_reason="evidence_found" if records else "no_relevant_evidence",
|
|
)
|
|
for page, _ in selected_pages:
|
|
page.touch()
|
|
self._refresh_tiers()
|
|
return records, decision
|
|
|
|
def correct(
|
|
self,
|
|
*,
|
|
text: str,
|
|
key: Tensor,
|
|
entity: str,
|
|
attribute: str,
|
|
value: str,
|
|
evidence: Optional[Sequence[str]] = None,
|
|
confidence: float = 1.0,
|
|
) -> tuple[MemoryRecordV2, str]:
|
|
"""Write an explicit correction as a new version."""
|
|
|
|
return self.write(
|
|
text=text,
|
|
key=key,
|
|
entity=entity,
|
|
attribute=attribute,
|
|
value=value,
|
|
evidence=evidence,
|
|
confidence=confidence,
|
|
importance=1.0,
|
|
source="user_correction",
|
|
trusted=True,
|
|
force=True,
|
|
)
|
|
|
|
def approve(self, record_id: str) -> MemoryRecordV2:
|
|
record = self.quarantine.pop(record_id, None)
|
|
if record is None:
|
|
raise KeyError(f"unknown quarantined record: {record_id}")
|
|
if self.tier_store is not None:
|
|
self.tier_store.delete_quarantine(record_id)
|
|
record.status = STATUS_ACTIVE
|
|
page = self._target_page(record.key, importance=record.importance)
|
|
self.records[record.record_id] = record
|
|
self._update_page(page, record)
|
|
self._store_record(record)
|
|
conflict_key = record.conflict_key()
|
|
if conflict_key:
|
|
old_id = self.active_by_conflict.get(conflict_key)
|
|
if old_id and old_id not in self.records:
|
|
self._hydrate_records([old_id])
|
|
old = self.records.get(old_id) if old_id else None
|
|
if old is not None and old.status == STATUS_ACTIVE:
|
|
old.status = STATUS_SUPERSEDED
|
|
record.version = old.version + 1
|
|
record.supersedes = old.record_id
|
|
self._store_record(old)
|
|
self.active_by_conflict[conflict_key] = record.record_id
|
|
self._refresh_tiers()
|
|
return record
|
|
|
|
def retract(self, record_id: str) -> None:
|
|
record = self.records.get(record_id)
|
|
if record is None and self.tier_store is not None:
|
|
self._hydrate_records([record_id])
|
|
record = self.records.get(record_id)
|
|
if record is None:
|
|
raise KeyError(record_id)
|
|
record.status = STATUS_RETRACTED
|
|
self._remove_gpu_record(record_id)
|
|
self._store_record(record)
|
|
conflict_key = record.conflict_key()
|
|
if conflict_key and self.active_by_conflict.get(conflict_key) == record_id:
|
|
del self.active_by_conflict[conflict_key]
|
|
self._refresh_tiers()
|
|
|
|
def consolidate(self, *, max_pages: Optional[int] = None) -> int:
|
|
"""Move low-value pages toward compact summaries without deleting evidence."""
|
|
|
|
target = max_pages or self.max_pages
|
|
if len(self.pages) <= target:
|
|
self._refresh_tiers()
|
|
return 0
|
|
ranked = sorted(self.pages.values(), key=lambda page: (page.importance, page.last_access))
|
|
merged = 0
|
|
for source in ranked:
|
|
if len(self.pages) <= target:
|
|
break
|
|
if source.tier == "hot" or not source.record_ids:
|
|
continue
|
|
destinations = [
|
|
page
|
|
for page in self.pages.values()
|
|
if page.page_id != source.page_id and not page.full
|
|
]
|
|
# Consolidation may only use an existing page. Creating a new
|
|
# page while trying to enforce a page limit would recurse forever
|
|
# when every page is already full.
|
|
if not destinations:
|
|
continue
|
|
source_key = source.key if source.key is not None else torch.zeros(self.key_dim)
|
|
source_key = self._encode_key(source_key)
|
|
destination = max(
|
|
destinations,
|
|
key=lambda page: (
|
|
float(torch.dot(source_key, F.normalize(page.key.float(), dim=0)))
|
|
if page.key is not None
|
|
else -0.1 * len(page.record_ids)
|
|
),
|
|
)
|
|
for record_id in list(source.record_ids):
|
|
if record_id not in destination.record_ids and not destination.full:
|
|
destination.record_ids.append(record_id)
|
|
self.records[record_id].page_id = destination.page_id
|
|
if destination.full:
|
|
self._open_page_ids.discard(destination.page_id)
|
|
else:
|
|
self._open_page_ids.add(destination.page_id)
|
|
self._remove_page_from_index(destination)
|
|
self._index_page(destination)
|
|
if destination.key is None:
|
|
destination.key = source.key
|
|
if destination.summary is None:
|
|
destination.summary = source.summary
|
|
destination.importance = max(destination.importance, source.importance)
|
|
self._remove_page_from_index(source)
|
|
del self.pages[source.page_id]
|
|
self._open_page_ids.discard(source.page_id)
|
|
merged += 1
|
|
self._refresh_tiers()
|
|
return merged
|
|
|
|
def stats(self) -> dict[str, Any]:
|
|
active = [record for record in self.records.values() if record.status == STATUS_ACTIVE]
|
|
stored = self.tier_store.count() if self.tier_store is not None else {}
|
|
record_count = int(stored.get("records", len(self.records)))
|
|
active_count = int(stored.get("status_active", len(active)))
|
|
superseded_count = int(
|
|
stored.get(
|
|
"status_superseded",
|
|
sum(record.status == STATUS_SUPERSEDED for record in self.records.values()),
|
|
)
|
|
)
|
|
retracted_count = int(
|
|
stored.get(
|
|
"status_retracted",
|
|
sum(record.status == STATUS_RETRACTED for record in self.records.values()),
|
|
)
|
|
)
|
|
quarantined_count = int(stored.get("quarantined", len(self.quarantine)))
|
|
return {
|
|
"pages": len(self.pages),
|
|
"max_pages": self.max_pages,
|
|
"page_capacity": self.page_capacity,
|
|
"capacity_records": self.max_pages * self.page_capacity,
|
|
"open_pages": len(self._open_page_ids),
|
|
"records": record_count,
|
|
"active_records": active_count,
|
|
"superseded_records": superseded_count,
|
|
"retracted_records": retracted_count,
|
|
"quarantined_records": quarantined_count,
|
|
"hot_pages": int(stored.get("hot_pages", sum(page.tier == "hot" for page in self.pages.values()))),
|
|
"warm_pages": int(stored.get("warm_pages", sum(page.tier == "warm" for page in self.pages.values()))),
|
|
"cold_pages": int(stored.get("cold_pages", sum(page.tier == "cold" for page in self.pages.values()))),
|
|
"active_conflict_keys": len(self.active_by_conflict),
|
|
"coarse_index_buckets": (
|
|
self.tier_store.coarse_bucket_count()
|
|
if self.tier_store is not None
|
|
else len(self._coarse_buckets)
|
|
),
|
|
"last_coarse_candidates": self._last_coarse_candidates,
|
|
"storage_mode": "tiered" if self.tier_store is not None else "embedded",
|
|
"resident_records": len(self.records),
|
|
"resident_pages": len(self.pages),
|
|
"gpu_cache_records": len(self._gpu_record_cache),
|
|
"gpu_cache_tokens": self._gpu_cache_tokens_used,
|
|
"gpu_cache_bytes": self._gpu_cache_bytes_used,
|
|
"gpu_cache_hits": self._gpu_cache_hits,
|
|
"gpu_cache_misses": self._gpu_cache_misses,
|
|
"gpu_cache_fallbacks": self._gpu_cache_fallbacks,
|
|
"gpu_cache_alloc_failures": self._gpu_cache_alloc_failures,
|
|
"gpu_cache_device": str(self._gpu_device) if self._gpu_device is not None else "none",
|
|
"gpu_cache_limit_records": self.gpu_cache_records,
|
|
"gpu_cache_limit_tokens": self.gpu_cache_tokens,
|
|
"gpu_cache_reserve_mb": self.gpu_cache_reserve_mb,
|
|
"gpu_cache_adaptive": self.gpu_cache_adaptive,
|
|
"gpu_cache_last_free_mb": (
|
|
round(self._gpu_cache_last_free_bytes / (1024 * 1024), 2)
|
|
if self._gpu_cache_last_free_bytes is not None
|
|
else None
|
|
),
|
|
}
|
|
|
|
def _resolve_record(self, record_id: str) -> MemoryRecordV2:
|
|
record = self.records.get(str(record_id))
|
|
if record is None and self.tier_store is not None:
|
|
self._hydrate_records([str(record_id)])
|
|
record = self.records.get(str(record_id))
|
|
if record is None:
|
|
record = self.quarantine.get(str(record_id))
|
|
if record is None:
|
|
raise KeyError(str(record_id))
|
|
return record
|
|
|
|
def list_records(
|
|
self,
|
|
*,
|
|
query_text: str = "",
|
|
status: str = STATUS_ACTIVE,
|
|
limit: int = 100,
|
|
offset: int = 0,
|
|
) -> list[MemoryRecordV2]:
|
|
"""List records for the local management API without exposing vectors."""
|
|
|
|
if limit < 1 or limit > 10000:
|
|
raise ValueError("limit must be between 1 and 10000")
|
|
if offset < 0:
|
|
raise ValueError("offset must be non-negative")
|
|
allowed = {STATUS_ACTIVE, STATUS_SUPERSEDED, STATUS_RETRACTED, STATUS_QUARANTINED, "all"}
|
|
if status not in allowed:
|
|
raise ValueError(f"status must be one of: {', '.join(sorted(allowed))}")
|
|
candidates = list(self.records.values()) + list(self.quarantine.values())
|
|
if status != "all":
|
|
candidates = [item for item in candidates if item.status == status]
|
|
if query_text.strip():
|
|
candidates = [
|
|
item for item in candidates
|
|
if item.lexical_score(query_text) > 0.0 or query_text.strip().lower() in item.text.lower()
|
|
]
|
|
candidates.sort(
|
|
key=lambda item: (item.lexical_score(query_text), item.last_access, item.timestamp),
|
|
reverse=True,
|
|
)
|
|
else:
|
|
candidates.sort(key=lambda item: (item.last_access, item.timestamp), reverse=True)
|
|
return candidates[offset : offset + limit]
|
|
|
|
def edit_record(
|
|
self,
|
|
record_id: str,
|
|
*,
|
|
text: Optional[str] = None,
|
|
key: Optional[Tensor] = None,
|
|
summary: Optional[Tensor] = None,
|
|
memory_type: Optional[str] = None,
|
|
entity: Optional[str] = None,
|
|
attribute: Optional[str] = None,
|
|
value: Optional[str] = None,
|
|
importance: Optional[float] = None,
|
|
confidence: Optional[float] = None,
|
|
evidence: Optional[Sequence[str]] = None,
|
|
token_ids: Optional[Tensor] = None,
|
|
token_mask: Optional[Tensor] = None,
|
|
source: str = "user_edit",
|
|
) -> MemoryRecordV2:
|
|
"""Create a new version and preserve the old record as superseded."""
|
|
|
|
old = self._resolve_record(record_id)
|
|
if old.status != STATUS_ACTIVE:
|
|
raise ValueError(f"only active records can be edited: {record_id}")
|
|
old.status = STATUS_SUPERSEDED
|
|
self._remove_gpu_record(old.record_id)
|
|
self._store_record(old)
|
|
old_conflict = old.conflict_key()
|
|
if old_conflict and self.active_by_conflict.get(old_conflict) == old.record_id:
|
|
del self.active_by_conflict[old_conflict]
|
|
next_text = old.text if text is None else str(text)
|
|
next_entity = old.entity if entity is None else str(entity)
|
|
next_attribute = old.attribute if attribute is None else str(attribute)
|
|
next_value = old.value if value is None else str(value)
|
|
next_importance = old.importance if importance is None else float(importance)
|
|
next_confidence = old.confidence if confidence is None else float(confidence)
|
|
next_key = old.key if key is None else key
|
|
next_summary = old.summary if summary is None else summary
|
|
now = _now()
|
|
new_record = MemoryRecordV2(
|
|
record_id=_stable_id(f"{old.record_id}|edit|{now}|{len(self.records)}", prefix="mem"),
|
|
text=next_text,
|
|
key=self._encode_key(next_key).cpu().clone(),
|
|
summary=self._encode_key(next_summary).cpu().clone(),
|
|
memory_type=old.memory_type if memory_type is None else str(memory_type),
|
|
entity=next_entity,
|
|
attribute=next_attribute,
|
|
value=next_value,
|
|
timestamp=now,
|
|
importance=max(0.0, min(1.0, next_importance)),
|
|
confidence=max(0.0, min(1.0, next_confidence)),
|
|
source=str(source),
|
|
status=STATUS_ACTIVE,
|
|
version=old.version + 1,
|
|
supersedes=old.record_id,
|
|
related_ids=list(old.related_ids),
|
|
evidence=list(old.evidence if evidence is None else evidence),
|
|
slot_index=old.slot_index,
|
|
token_ids=(token_ids.detach().cpu().long().reshape(-1).clone() if token_ids is not None else old.token_ids),
|
|
token_mask=(token_mask.detach().cpu().bool().reshape(-1).clone() if token_mask is not None else old.token_mask),
|
|
)
|
|
self.records[new_record.record_id] = new_record
|
|
if new_record.text.strip():
|
|
self._text_index[new_record.text.strip().lower()] = new_record.record_id
|
|
page = self._target_page(new_record.key, importance=new_record.importance)
|
|
self._update_page(page, new_record)
|
|
self._store_record(new_record)
|
|
new_conflict = new_record.conflict_key()
|
|
if new_conflict:
|
|
self.active_by_conflict[new_conflict] = new_record.record_id
|
|
self._refresh_tiers()
|
|
return new_record
|
|
|
|
def _remove_gpu_record(self, record_id: str) -> None:
|
|
cached = self._gpu_record_cache.pop(record_id, None)
|
|
if cached is None:
|
|
return
|
|
if record_id in self._gpu_cache_order:
|
|
self._gpu_cache_order.remove(record_id)
|
|
token_ids = cached.get("token_ids")
|
|
if isinstance(token_ids, Tensor):
|
|
self._gpu_cache_tokens_used -= int(token_ids.numel())
|
|
self._gpu_cache_bytes_used -= sum(self._tensor_bytes(value) for value in cached.values())
|
|
self._gpu_cache_tokens_used = max(0, self._gpu_cache_tokens_used)
|
|
self._gpu_cache_bytes_used = max(0, self._gpu_cache_bytes_used)
|
|
|
|
def audit(self) -> dict[str, Any]:
|
|
"""Check page membership, conflict indexes and relationship references."""
|
|
|
|
issues: list[str] = []
|
|
page_membership: set[str] = set()
|
|
for page in self.pages.values():
|
|
if len(page.record_ids) > page.capacity:
|
|
issues.append(f"page_over_capacity:{page.page_id}")
|
|
for record_id in page.record_ids:
|
|
page_membership.add(record_id)
|
|
record = self.records.get(record_id)
|
|
if record is None:
|
|
issues.append(f"missing_record:{record_id}")
|
|
elif record.page_id != page.page_id:
|
|
issues.append(f"page_pointer_mismatch:{record_id}")
|
|
for record in self.records.values():
|
|
if record.status == STATUS_ACTIVE and record.record_id not in page_membership:
|
|
issues.append(f"active_record_not_indexed:{record.record_id}")
|
|
for related_id in record.related_ids:
|
|
if related_id not in self.records and related_id not in self.quarantine:
|
|
issues.append(f"dangling_related_id:{record.record_id}->{related_id}")
|
|
for conflict_key, record_id in self.active_by_conflict.items():
|
|
record = self.records.get(record_id)
|
|
if record is None or record.status != STATUS_ACTIVE or record.conflict_key() != conflict_key:
|
|
issues.append(f"invalid_conflict_index:{conflict_key}")
|
|
return {
|
|
"healthy": not issues,
|
|
"issues": issues[:100],
|
|
"issue_count": len(issues),
|
|
"stats": self.stats(),
|
|
}
|
|
|
|
def flush_storage(self) -> None:
|
|
"""Flush the durable tier without serializing cold pages into RAM."""
|
|
|
|
if self.tier_store is not None:
|
|
self.tier_store.flush()
|
|
|
|
def close_storage(self) -> None:
|
|
"""Close the durable page database so Windows can remove or rotate it."""
|
|
|
|
if self.tier_store is not None:
|
|
self.tier_store.close()
|
|
|
|
def export_payload(self) -> dict[str, Any]:
|
|
"""Export tensors plus JSON-safe metadata for checkpoint backends."""
|
|
|
|
if self.tier_store is not None:
|
|
for page in self.pages.values():
|
|
self._hydrate_page(page.page_id)
|
|
records = []
|
|
for record in self.records.values():
|
|
item = asdict(record)
|
|
item["key"] = record.key.detach().cpu()
|
|
item["summary"] = record.summary.detach().cpu()
|
|
records.append(item)
|
|
pages = []
|
|
for page in self.pages.values():
|
|
item = asdict(page)
|
|
item["key"] = page.key.detach().cpu() if page.key is not None else None
|
|
item["summary"] = page.summary.detach().cpu() if page.summary is not None else None
|
|
pages.append(item)
|
|
quarantine = []
|
|
for record in self.quarantine.values():
|
|
item = asdict(record)
|
|
item["key"] = record.key.detach().cpu()
|
|
item["summary"] = record.summary.detach().cpu()
|
|
quarantine.append(item)
|
|
return {
|
|
"format_version": 2,
|
|
"hidden_size": self.hidden_size,
|
|
"page_capacity": self.page_capacity,
|
|
"max_pages": self.max_pages,
|
|
"hot_pages": self.hot_pages,
|
|
"top_k_pages": self.top_k_pages,
|
|
"top_k_records": self.top_k_records,
|
|
"max_hops": self.max_hops,
|
|
"key_dim": self.key_dim,
|
|
"coarse_index_bits": self.coarse_index_bits,
|
|
"record_count": len(self.records),
|
|
"counter": self._counter,
|
|
"records": records,
|
|
"pages": pages,
|
|
"quarantine": quarantine,
|
|
"active_by_conflict": dict(self.active_by_conflict),
|
|
}
|
|
|
|
@classmethod
|
|
def from_payload(
|
|
cls,
|
|
payload: dict[str, Any],
|
|
*,
|
|
router: Optional[MemoryRouterV2] = None,
|
|
tier_store: Optional["TieredMemoryStoreV2"] = None,
|
|
max_resident_pages: int = 256,
|
|
runtime_device: Optional[torch.device] = None,
|
|
gpu_cache_records: int = 256,
|
|
gpu_cache_tokens: int = 131072,
|
|
gpu_cache_reserve_mb: int = 2048,
|
|
gpu_cache_adaptive: bool = True,
|
|
) -> "PagedMemoryBankV2":
|
|
bank = cls(
|
|
int(payload["hidden_size"]),
|
|
page_capacity=int(payload.get("page_capacity", 32)),
|
|
max_pages=int(payload.get("max_pages", 32768)),
|
|
hot_pages=int(payload.get("hot_pages", 8)),
|
|
top_k_pages=int(payload.get("top_k_pages", 4)),
|
|
top_k_records=int(payload.get("top_k_records", 8)),
|
|
max_hops=int(payload.get("max_hops", 3)),
|
|
router=router,
|
|
key_dim=int(payload.get("key_dim", router.router_dim if router is not None else payload["hidden_size"])),
|
|
coarse_index_bits=int(payload.get("coarse_index_bits", 20)),
|
|
tier_store=tier_store,
|
|
max_resident_pages=max_resident_pages,
|
|
runtime_device=runtime_device,
|
|
gpu_cache_records=gpu_cache_records,
|
|
gpu_cache_tokens=gpu_cache_tokens,
|
|
gpu_cache_reserve_mb=gpu_cache_reserve_mb,
|
|
gpu_cache_adaptive=gpu_cache_adaptive,
|
|
)
|
|
bank._counter = int(payload.get("counter", 0))
|
|
for item in payload.get("records", []):
|
|
item = dict(item)
|
|
item["key"] = torch.as_tensor(item["key"]).float()
|
|
item["summary"] = torch.as_tensor(item["summary"]).float()
|
|
if item.get("token_ids") is not None:
|
|
item["token_ids"] = torch.as_tensor(item["token_ids"]).long()
|
|
if item.get("token_mask") is not None:
|
|
item["token_mask"] = torch.as_tensor(item["token_mask"]).bool()
|
|
bank.records[item["record_id"]] = MemoryRecordV2(**item)
|
|
if item.get("text", "").strip() and item.get("status") == STATUS_ACTIVE:
|
|
bank._text_index[item["text"].strip().lower()] = item["record_id"]
|
|
for item in payload.get("pages", []):
|
|
item = dict(item)
|
|
if item.get("key") is not None:
|
|
item["key"] = torch.as_tensor(item["key"]).float()
|
|
if item.get("summary") is not None:
|
|
item["summary"] = torch.as_tensor(item["summary"]).float()
|
|
bank.pages[item["page_id"]] = MemoryPageV2(**item)
|
|
bank._open_page_ids = {
|
|
page.page_id for page in bank.pages.values() if not page.full
|
|
}
|
|
bank._open_page_order = [
|
|
page.page_id
|
|
for page in sorted(
|
|
bank.pages.values(), key=lambda item: (item.created_at, item.page_id)
|
|
)
|
|
if page.page_id in bank._open_page_ids
|
|
]
|
|
for page in bank.pages.values():
|
|
bank._index_page(page)
|
|
for item in payload.get("quarantine", []):
|
|
item = dict(item)
|
|
item["key"] = torch.as_tensor(item["key"]).float()
|
|
item["summary"] = torch.as_tensor(item["summary"]).float()
|
|
if item.get("token_ids") is not None:
|
|
item["token_ids"] = torch.as_tensor(item["token_ids"]).long()
|
|
if item.get("token_mask") is not None:
|
|
item["token_mask"] = torch.as_tensor(item["token_mask"]).bool()
|
|
bank.quarantine[item["record_id"]] = MemoryRecordV2(**item)
|
|
bank.active_by_conflict = dict(payload.get("active_by_conflict", {}))
|
|
bank._refresh_tiers()
|
|
return bank
|
|
|
|
|
|
@dataclass
|
|
class KVBudgetManagerV2:
|
|
"""Token budget and overflow policy for the hot working context."""
|
|
|
|
max_tokens: int = 32768
|
|
hard_max_tokens: int = 131072
|
|
compaction_trigger: float = 0.90
|
|
keep_recent_tokens: int = 8192
|
|
|
|
def __post_init__(self) -> None:
|
|
if not 1 <= self.max_tokens <= self.hard_max_tokens:
|
|
raise ValueError("max_tokens must be inside [1, hard_max_tokens]")
|
|
if not 0.5 <= self.compaction_trigger < 1.0:
|
|
raise ValueError("compaction_trigger must be in [0.5, 1)")
|
|
self.keep_recent_tokens = min(self.keep_recent_tokens, self.max_tokens)
|
|
|
|
@property
|
|
def trigger_tokens(self) -> int:
|
|
return max(1, int(self.max_tokens * self.compaction_trigger))
|
|
|
|
def needs_compaction(self, token_count: int) -> bool:
|
|
return int(token_count) >= self.trigger_tokens
|
|
|
|
def overflow(self, token_count: int) -> int:
|
|
return max(0, int(token_count) - self.max_tokens)
|
|
|
|
def retention_plan(self, token_count: int) -> dict[str, int | bool]:
|
|
overflow = self.overflow(token_count)
|
|
return {
|
|
"token_count": int(token_count),
|
|
"max_tokens": self.max_tokens,
|
|
"needs_compaction": self.needs_compaction(token_count),
|
|
"overflow_tokens": overflow,
|
|
"preserve_recent_tokens": self.keep_recent_tokens,
|
|
"tokens_for_memory": max(0, overflow + max(0, int(token_count * 0.1))),
|
|
}
|
|
|
|
|
|
class MemoryOSV2:
|
|
"""Model-facing coordinator for a hierarchical memory bank."""
|
|
|
|
def __init__(
|
|
self,
|
|
hidden_size: int,
|
|
*,
|
|
router: Optional[MemoryRouterV2] = None,
|
|
bank: Optional[PagedMemoryBankV2] = None,
|
|
kv_budget: Optional[KVBudgetManagerV2] = None,
|
|
read_threshold: float = 0.65,
|
|
write_threshold: float = 0.50,
|
|
runtime_device: Optional[torch.device] = None,
|
|
gpu_cache_records: int = 256,
|
|
gpu_cache_tokens: int = 131072,
|
|
gpu_cache_reserve_mb: int = 2048,
|
|
gpu_cache_adaptive: bool = True,
|
|
) -> None:
|
|
self.router = router or MemoryRouterV2(hidden_size)
|
|
self.bank = bank or PagedMemoryBankV2(
|
|
hidden_size,
|
|
router=self.router,
|
|
runtime_device=runtime_device,
|
|
gpu_cache_records=gpu_cache_records,
|
|
gpu_cache_tokens=gpu_cache_tokens,
|
|
gpu_cache_reserve_mb=gpu_cache_reserve_mb,
|
|
gpu_cache_adaptive=gpu_cache_adaptive,
|
|
)
|
|
self.kv_budget = kv_budget or KVBudgetManagerV2()
|
|
self.read_threshold = read_threshold
|
|
self.write_threshold = write_threshold
|
|
|
|
def write(self, **kwargs: Any) -> tuple[MemoryRecordV2, str]:
|
|
importance = float(kwargs.get("importance", 0.5))
|
|
confidence = float(kwargs.get("confidence", 0.5))
|
|
force = bool(kwargs.get("force", False))
|
|
trusted = bool(kwargs.get("trusted", True))
|
|
if not force and (importance < self.write_threshold or confidence < 0.25):
|
|
kwargs["trusted"] = False
|
|
return self.bank.write(**kwargs)
|
|
|
|
def write_batch(self, records: Iterable[dict[str, Any]]) -> list[tuple[MemoryRecordV2, str]]:
|
|
"""Apply the V2 write gate and commit a batch atomically when possible."""
|
|
|
|
prepared: list[dict[str, Any]] = []
|
|
for item in records:
|
|
kwargs = dict(item)
|
|
importance = float(kwargs.get("importance", 0.5))
|
|
confidence = float(kwargs.get("confidence", 0.5))
|
|
force = bool(kwargs.get("force", False))
|
|
if not force and (importance < self.write_threshold or confidence < 0.25):
|
|
kwargs["trusted"] = False
|
|
prepared.append(kwargs)
|
|
return self.bank.write_batch(prepared)
|
|
|
|
def read(
|
|
self,
|
|
*,
|
|
query_key: Tensor,
|
|
query_text: str = "",
|
|
query_token_ids: Optional[Tensor] = None,
|
|
top_k_pages: Optional[int] = None,
|
|
top_k_records: Optional[int] = None,
|
|
max_hops: Optional[int] = None,
|
|
) -> tuple[list[MemoryRecordV2], RouterDecisionV2]:
|
|
if self.router is not None and query_key.numel() == self.router.hidden_size:
|
|
router_device = next(self.router.parameters()).device
|
|
need_probability = float(
|
|
torch.sigmoid(self.router.need_memory(query_key.to(router_device).reshape(1, -1)))[0, 0].item()
|
|
)
|
|
token_evidence = self.bank.has_token_evidence(query_key, query_token_ids)
|
|
if need_probability < self.read_threshold and not token_evidence:
|
|
return [], RouterDecisionV2(
|
|
need_memory=False,
|
|
page_ids=[],
|
|
record_ids=[],
|
|
page_scores=[],
|
|
record_scores=[],
|
|
hop_count=0,
|
|
hop_trace=[],
|
|
confidence=need_probability,
|
|
stop_reason="router_abstained",
|
|
)
|
|
records, decision = self.bank.query(
|
|
query_key=query_key,
|
|
query_text=query_text,
|
|
query_token_ids=query_token_ids,
|
|
top_k_pages=top_k_pages,
|
|
top_k_records=top_k_records,
|
|
max_hops=max_hops,
|
|
min_score=-1.0,
|
|
)
|
|
token_evidence = self.bank.has_token_evidence(query_key, query_token_ids)
|
|
if decision.confidence < self.read_threshold and not token_evidence:
|
|
decision.need_memory = False
|
|
decision.record_ids = []
|
|
decision.stop_reason = "below_read_threshold"
|
|
return [], decision
|
|
if token_evidence and decision.confidence < self.read_threshold:
|
|
decision.stop_reason = "token_evidence_override"
|
|
return records, decision
|
|
|
|
def correct(self, **kwargs: Any) -> tuple[MemoryRecordV2, str]:
|
|
return self.bank.correct(**kwargs)
|
|
|
|
def approve(self, record_id: str) -> MemoryRecordV2:
|
|
return self.bank.approve(record_id)
|
|
|
|
def retract(self, record_id: str) -> None:
|
|
self.bank.retract(record_id)
|
|
|
|
def stats(self) -> dict[str, Any]:
|
|
output = self.bank.stats()
|
|
output["kv_budget"] = {
|
|
"max_tokens": self.kv_budget.max_tokens,
|
|
"trigger_tokens": self.kv_budget.trigger_tokens,
|
|
}
|
|
return output
|
|
|
|
def list_records(self, **kwargs: Any) -> list[MemoryRecordV2]:
|
|
return self.bank.list_records(**kwargs)
|
|
|
|
def edit_record(self, record_id: str, **kwargs: Any) -> MemoryRecordV2:
|
|
return self.bank.edit_record(record_id, **kwargs)
|
|
|
|
def retract_record(self, record_id: str) -> None:
|
|
self.bank.retract(record_id)
|
|
|
|
def audit(self) -> dict[str, Any]:
|
|
return self.bank.audit()
|
|
|
|
def flush_storage(self) -> None:
|
|
self.bank.flush_storage()
|
|
|
|
def close_storage(self) -> None:
|
|
self.bank.close_storage()
|
|
|
|
def export_payload(self) -> dict[str, Any]:
|
|
payload = self.bank.export_payload()
|
|
payload["read_threshold"] = self.read_threshold
|
|
payload["write_threshold"] = self.write_threshold
|
|
payload["kv_budget"] = asdict(self.kv_budget)
|
|
return payload
|
|
|
|
@classmethod
|
|
def from_payload(
|
|
cls,
|
|
payload: dict[str, Any],
|
|
*,
|
|
router: Optional[MemoryRouterV2] = None,
|
|
tier_store: Optional["TieredMemoryStoreV2"] = None,
|
|
max_resident_pages: int = 256,
|
|
runtime_device: Optional[torch.device] = None,
|
|
gpu_cache_records: int = 256,
|
|
gpu_cache_tokens: int = 131072,
|
|
gpu_cache_reserve_mb: int = 2048,
|
|
gpu_cache_adaptive: bool = True,
|
|
) -> "MemoryOSV2":
|
|
bank = PagedMemoryBankV2.from_payload(
|
|
payload,
|
|
router=router,
|
|
tier_store=tier_store,
|
|
max_resident_pages=max_resident_pages,
|
|
runtime_device=runtime_device,
|
|
gpu_cache_records=gpu_cache_records,
|
|
gpu_cache_tokens=gpu_cache_tokens,
|
|
gpu_cache_reserve_mb=gpu_cache_reserve_mb,
|
|
gpu_cache_adaptive=gpu_cache_adaptive,
|
|
)
|
|
budget = KVBudgetManagerV2(**payload.get("kv_budget", {}))
|
|
return cls(
|
|
bank.hidden_size,
|
|
router=router,
|
|
bank=bank,
|
|
kv_budget=budget,
|
|
read_threshold=float(payload.get("read_threshold", 0.65)),
|
|
write_threshold=float(payload.get("write_threshold", 0.5)),
|
|
)
|
|
|
|
|
|
def router_training_loss(
|
|
router: MemoryRouterV2,
|
|
query: Tensor,
|
|
candidates: Tensor,
|
|
positive_index: Tensor,
|
|
*,
|
|
need_memory_label: Optional[Tensor] = None,
|
|
hop_label: Optional[Tensor] = None,
|
|
) -> dict[str, Tensor]:
|
|
"""Compute supervised router losses with hard negatives."""
|
|
|
|
output = router(query, candidates)
|
|
losses: dict[str, Tensor] = {}
|
|
losses["candidate"] = F.cross_entropy(output["scores"], positive_index)
|
|
if need_memory_label is not None:
|
|
losses["need_memory"] = F.binary_cross_entropy_with_logits(
|
|
output["need_memory_logits"], need_memory_label.float()
|
|
)
|
|
if hop_label is not None:
|
|
losses["hop"] = F.cross_entropy(output["hop_logits"], hop_label)
|
|
losses["total"] = sum(losses.values())
|
|
return losses
|
|
|
|
|
|
__all__ = [
|
|
"KVBudgetManagerV2",
|
|
"MemoryOSV2",
|
|
"MemoryPageV2",
|
|
"MemoryRecordV2",
|
|
"MemoryRouterV2",
|
|
"PagedMemoryBankV2",
|
|
"RouterDecisionV2",
|
|
"router_training_loss",
|
|
"STATUS_ACTIVE",
|
|
"STATUS_QUARANTINED",
|
|
"STATUS_RETRACTED",
|
|
"STATUS_SUPERSEDED",
|
|
]
|