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