Files

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",
]