Add Natural Memory architecture and tooling

This commit is contained in:
WpyQwq
2026-09-05 08:53:41 +08:00
parent 0acf8b06ee
commit 516351f0b5
56 changed files with 18319 additions and 0 deletions
View File
+309
View File
@@ -0,0 +1,309 @@
from __future__ import annotations
import unittest
from tempfile import TemporaryDirectory
import torch
from dynamic_memory_lab.memory_os_v2 import (
KVBudgetManagerV2,
MemoryOSV2,
MemoryRouterV2,
PagedMemoryBankV2,
STATUS_ACTIVE,
STATUS_QUARANTINED,
STATUS_SUPERSEDED,
)
from dynamic_memory_lab.tiered_memory_store_v2 import TieredMemoryStoreV2
class MemoryOSV2Test(unittest.TestCase):
def setUp(self) -> None:
torch.manual_seed(7)
self.router = MemoryRouterV2(16, router_dim=8, num_heads=2, max_hops=3)
self.bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
max_pages=512,
hot_pages=2,
top_k_pages=2,
top_k_records=4,
max_hops=3,
coarse_index_bits=8,
)
def test_router_scores_and_compressed_address(self) -> None:
query = torch.randn(4, 16)
candidates = torch.randn(4, 5, 16)
output = self.router(query, candidates)
self.assertEqual(tuple(output["scores"].shape), (4, 5))
self.assertEqual(tuple(output["head_scores"].shape), (4, 5, 2))
self.assertEqual(tuple(self.router.encode_key(query).shape), (4, 8))
def test_write_version_and_conflict_resolution(self) -> None:
first, first_action = self.bank.write(
text="我住在上海",
key=torch.randn(16),
entity="user",
attribute="city",
value="上海",
confidence=0.9,
)
second, second_action = self.bank.write(
text="我搬到了杭州",
key=torch.randn(16),
entity="user",
attribute="city",
value="杭州",
confidence=0.95,
)
self.assertEqual(first_action, "inserted")
self.assertEqual(second_action, "updated")
self.assertEqual(first.status, STATUS_SUPERSEDED)
self.assertEqual(second.status, STATUS_ACTIVE)
self.assertEqual(second.version, 1)
self.assertEqual(self.bank.active_by_conflict["user::city"], second.record_id)
def test_quarantine_and_approval(self) -> None:
record, action = self.bank.write(
text="未经确认的推断",
key=torch.randn(16),
trusted=False,
confidence=0.1,
)
self.assertEqual(action, "quarantined")
self.assertEqual(record.status, STATUS_QUARANTINED)
self.assertNotIn(record.record_id, self.bank.records)
approved = self.bank.approve(record.record_id)
self.assertEqual(approved.status, STATUS_ACTIVE)
self.assertIn(approved.record_id, self.bank.records)
def test_multi_hop_and_slot_replacement(self) -> None:
second, _ = self.bank.write(text="项目的第二个节点", key=torch.randn(16), slot_index=2)
third, _ = self.bank.write(text="项目的第三个节点", key=torch.randn(16), slot_index=3)
first, _ = self.bank.write(
text="项目的第一个节点",
key=torch.randn(16),
related_ids=[second.record_id, third.record_id],
slot_index=1,
)
replacement, action = self.bank.write(
text="项目的第一个节点修正版",
key=torch.randn(16),
related_ids=[second.record_id],
slot_index=1,
)
self.assertEqual(action, "updated")
self.assertEqual(first.status, STATUS_SUPERSEDED)
self.assertEqual(replacement.status, STATUS_ACTIVE)
records, decision = self.bank.query(
query_key=replacement.key,
top_k_pages=2,
top_k_records=4,
max_hops=3,
)
ids = {record.record_id for record in records}
self.assertIn(replacement.record_id, ids)
self.assertGreaterEqual(decision.hop_count, 1)
def test_coarse_index_bounds_candidate_pages(self) -> None:
for index in range(300):
key = torch.zeros(16)
key[index % 16] = 1.0
key[(index * 7 + 3) % 16] += 0.05
self.bank.write(text=f"memory-{index}", key=key, importance=0.2)
query_key = torch.zeros(16)
query_key[3] = 1.0
self.bank.query(query_key=query_key, top_k_pages=2, top_k_records=2)
stats = self.bank.stats()
self.assertGreater(stats["pages"], 128)
self.assertLess(stats["last_coarse_candidates"], stats["pages"])
def test_export_and_restore(self) -> None:
record, _ = self.bank.write(
text="可持久化事实",
key=torch.randn(16),
token_ids=torch.tensor([4, 5, 6]),
token_mask=torch.tensor([True, True, True]),
)
payload = self.bank.export_payload()
restored = PagedMemoryBankV2.from_payload(payload, router=self.router)
self.assertEqual(restored.stats()["active_records"], 1)
self.assertTrue(torch.equal(restored.records[record.record_id].token_ids, torch.tensor([4, 5, 6])))
self.assertEqual(restored.records[record.record_id].page_id, record.page_id)
def test_lazy_capacity_is_bounded(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
max_pages=2,
hot_pages=0,
coarse_index_bits=8,
)
bank.write(text="容量一", key=torch.randn(16))
bank.write(text="容量二", key=torch.randn(16))
self.assertEqual(bank.stats()["pages"], 2)
with self.assertRaises(RuntimeError):
bank.write(text="容量三", key=torch.randn(16))
def test_memory_os_and_kv_budget(self) -> None:
os_v2 = MemoryOSV2(16, router=self.router)
record, action = os_v2.write(
text="可靠事实",
key=torch.randn(16),
importance=0.9,
confidence=0.9,
)
self.assertEqual(action, "inserted")
self.assertIn(record.record_id, os_v2.bank.records)
budget = KVBudgetManagerV2(max_tokens=128, hard_max_tokens=512, keep_recent_tokens=32)
self.assertFalse(budget.needs_compaction(100))
self.assertTrue(budget.needs_compaction(120))
self.assertEqual(budget.overflow(140), 12)
def test_batch_context_records_keep_all_chunks_active(self) -> None:
os_v2 = MemoryOSV2(16, router=self.router)
output = os_v2.write_batch(
[
{
"text": "context_chunk:0:0:4",
"key": torch.randn(16),
"memory_type": "context_chunk",
"importance": 0.55,
"confidence": 0.8,
"trusted": True,
"force": True,
},
{
"text": "context_chunk:0:4:8",
"key": torch.randn(16),
"memory_type": "context_chunk",
"importance": 0.55,
"confidence": 0.8,
"trusted": True,
"force": True,
},
]
)
self.assertEqual(len(output), 2)
self.assertEqual(os_v2.stats()["active_records"], 2)
def test_management_list_edit_retract_and_audit(self) -> None:
os_v2 = MemoryOSV2(16, router=self.router)
record, _ = os_v2.write(
text="用户喜欢蓝色",
key=torch.randn(16),
entity="user",
attribute="color",
value="蓝色",
confidence=0.95,
importance=0.9,
)
listed = os_v2.list_records(query_text="蓝色", status="active", limit=10)
self.assertEqual([item.record_id for item in listed], [record.record_id])
edited = os_v2.edit_record(
record.record_id,
text="用户喜欢绿色",
entity="user",
attribute="color",
value="绿色",
evidence=["user_correction"],
)
self.assertEqual(edited.version, 1)
self.assertEqual(edited.supersedes, record.record_id)
self.assertEqual(os_v2.bank.records[record.record_id].status, STATUS_SUPERSEDED)
self.assertEqual(os_v2.list_records(query_text="绿色")[0].record_id, edited.record_id)
os_v2.retract_record(edited.record_id)
self.assertEqual(os_v2.bank.records[edited.record_id].status, "retracted")
audit = os_v2.audit()
self.assertTrue(audit["healthy"], audit)
all_records = os_v2.list_records(status="all", limit=10)
self.assertEqual(len(all_records), 2)
def test_tiered_storage_restarts_and_evicts_cold_records(self) -> None:
with TemporaryDirectory() as directory:
path = f"{directory}/memory.sqlite"
store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2)
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
max_pages=64,
hot_pages=1,
top_k_pages=2,
top_k_records=2,
tier_store=store,
max_resident_pages=1,
coarse_index_bits=8,
)
for index in range(8):
key = torch.zeros(16)
key[index % 8] = 1.0
bank.write(
text=f"tiered-memory-{index}",
key=key,
entity="user",
attribute=f"attr-{index}",
value=f"value-{index}",
importance=0.1 if index < 7 else 1.0,
confidence=0.95,
)
stats = bank.stats()
self.assertEqual(stats["storage_mode"], "tiered")
self.assertGreaterEqual(stats["pages"], 4)
self.assertGreater(stats["cold_pages"], 0)
self.assertLess(stats["resident_records"], stats["records"])
store.close()
reopened_store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2)
reopened = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
max_pages=64,
hot_pages=1,
top_k_pages=2,
top_k_records=2,
tier_store=reopened_store,
max_resident_pages=1,
coarse_index_bits=8,
)
records, decision = reopened.query(
query_key=torch.nn.functional.one_hot(torch.tensor(3), num_classes=16).float(),
query_text="tiered-memory-3",
top_k_pages=2,
top_k_records=2,
)
self.assertTrue(records)
self.assertTrue(any(item.text == "tiered-memory-3" for item in records))
self.assertGreaterEqual(decision.hop_count, 1)
quarantined, action = reopened.write(
text="待审批事实",
key=torch.randn(16),
trusted=False,
confidence=0.1,
)
self.assertEqual(action, "quarantined")
reopened_store.close()
final_store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2)
final_bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
max_pages=64,
hot_pages=1,
tier_store=final_store,
max_resident_pages=2,
coarse_index_bits=8,
)
self.assertIn(quarantined.record_id, final_bank.quarantine)
approved = final_bank.approve(quarantined.record_id)
self.assertEqual(approved.status, STATUS_ACTIVE)
final_store.close()
if __name__ == "__main__":
unittest.main()
+34
View File
@@ -0,0 +1,34 @@
from __future__ import annotations
import unittest
import torch
from dynamic_memory_lab.model import DynamicMemoryConfig, DynamicMemoryLM
from dynamic_memory_lab.tasks import sample_associative_batch
class DynamicMemoryModelTest(unittest.TestCase):
def setUp(self) -> None:
self.device = torch.device("cpu")
self.config = DynamicMemoryConfig(vocab_size=32, max_seq_len=16, d_model=32, n_layers=1, n_heads=4, memory_slots=2)
self.model = DynamicMemoryLM(self.config).to(self.device)
def test_shapes_and_loss(self) -> None:
batch = sample_associative_batch(batch_size=3, vocab_size=self.config.vocab_size, device=self.device)
memory = self.model(batch.learn_chunks[0]).memory
output = self.model(batch.query_input, memory=memory, update_memory=False, labels=batch.query_labels)
self.assertEqual(tuple(output.logits.shape), (3, 2, self.config.vocab_size))
self.assertEqual(tuple(output.memory.shape), (3, self.config.memory_slots, self.config.d_model))
self.assertIsNotNone(output.loss)
output.loss.backward()
def test_memory_changes_after_learning_chunk(self) -> None:
batch = sample_associative_batch(batch_size=2, vocab_size=self.config.vocab_size, device=self.device)
initial = self.model.memory.initial_state(2, device=self.device, dtype=torch.float32)
updated = self.model(batch.learn_chunks[0], memory=initial, update_memory=True).memory
self.assertFalse(torch.allclose(initial, updated))
if __name__ == "__main__":
unittest.main()
+142
View File
@@ -0,0 +1,142 @@
from __future__ import annotations
import unittest
import torch
from torch import nn
from dynamic_memory_lab.qwen_integration import (
AutomaticMemoryPolicy,
MemoryLayerAdapter,
NaturalLanguageRetriever,
NativeQwenDynamicMemory,
QwenMemoryConfig,
QwenDynamicMemory,
_MemoryRuntime,
looks_like_question,
split_memory_candidates,
)
class _FakeAttention(nn.Module):
def __init__(self, *, fail_if_called: bool = False) -> None:
super().__init__()
self.fail_if_called = fail_if_called
self.called = False
def forward(self, hidden_states: torch.Tensor, **kwargs):
self.called = True
if self.fail_if_called:
raise AssertionError("original token mixer was called in replace mode")
return hidden_states * 2.0, None
class _FakeQwenLayer(nn.Module):
layer_type = "full_attention"
def __init__(self, *, fail_if_called: bool = False) -> None:
super().__init__()
self.input_layernorm = nn.Identity()
self.post_attention_layernorm = nn.Identity()
self.self_attn = _FakeAttention(fail_if_called=fail_if_called)
self.mlp = nn.Identity()
def forward(self, hidden_states, position_embeddings=None, attention_mask=None, position_ids=None, past_key_values=None, **kwargs):
return hidden_states + self.self_attn(hidden_states)[0]
class QwenSurgeryTest(unittest.TestCase):
def _runtime(self):
memory = QwenDynamicMemory(
hidden_size=8,
config=QwenMemoryConfig(memory_slots=2, memory_dim=4),
)
runtime = _MemoryRuntime(memory)
runtime.state = memory.initial_state(1, device=torch.device("cpu"))
return runtime
def test_blend_exposes_trainable_mixer_weight(self) -> None:
runtime = self._runtime()
adapter = MemoryLayerAdapter(
_FakeQwenLayer(),
runtime,
read=True,
write=False,
mode="blend",
blend_init=0.5,
)
output = adapter(torch.ones(1, 3, 8))
output.sum().backward()
self.assertEqual(tuple(output.shape), (1, 3, 8))
self.assertIsNotNone(adapter.blend_logit.grad)
def test_replace_skips_original_token_mixer(self) -> None:
runtime = self._runtime()
layer = _FakeQwenLayer(fail_if_called=True)
adapter = MemoryLayerAdapter(layer, runtime, read=True, write=False, mode="replace")
output = adapter(torch.ones(1, 3, 8))
self.assertEqual(tuple(output.shape), (1, 3, 8))
self.assertFalse(layer.self_attn.called)
def test_raw_token_write_uses_output_projection_row(self) -> None:
memory = QwenDynamicMemory(
hidden_size=8,
config=QwenMemoryConfig(
memory_slots=2,
memory_dim=4,
write_token_offset=2,
raw_token_write=True,
broadcast_write=True,
),
)
runtime = _MemoryRuntime(memory)
runtime.state = memory.initial_state(1, device=torch.device("cpu"))
runtime.read_enabled = False
runtime.update_enabled = True
runtime.input_ids = torch.tensor([[5, 6, 7, 8]])
runtime.attention_mask = torch.ones_like(runtime.input_ids)
runtime.output_embeddings = nn.Linear(8, 16, bias=False)
adapter = MemoryLayerAdapter(_FakeQwenLayer(), runtime, read=True, write=True, mode="residual")
adapter(torch.ones(1, 4, 8))
expected = runtime.output_embeddings.weight[7]
self.assertIsNotNone(runtime.raw_memory)
self.assertTrue(torch.allclose(runtime.raw_memory[0], expected))
def test_native_controller_exposes_write_forget_and_value_state(self) -> None:
memory = NativeQwenDynamicMemory(
hidden_size=8,
config=QwenMemoryConfig(memory_slots=2, memory_dim=4),
)
hidden = torch.randn(1, 3, 8)
state = memory.initial_state(1, device=torch.device("cpu"))
updated = memory.update(hidden, state, attention_mask=torch.ones(1, 5, dtype=torch.long))
self.assertEqual(tuple(updated.shape), (1, 2, 4))
self.assertEqual(tuple(memory.last_write_probability.shape), (1, 1))
self.assertEqual(tuple(memory.last_forget_probability.shape), (1, 2))
self.assertEqual(tuple(memory.last_write_summary.shape), (1, 8))
self.assertEqual(tuple(memory.last_write_representation.shape), (1, 8))
def test_natural_language_retriever_scores_single_and_multiple_slots(self) -> None:
retriever = NaturalLanguageRetriever(hidden_size=8, projection_size=4)
query = torch.randn(2, 8)
one_key = torch.randn(2, 8)
many_keys = torch.randn(2, 3, 8)
self.assertEqual(tuple(retriever(query, one_key).shape), (2,))
self.assertEqual(tuple(retriever(query, many_keys).shape), (2, 3))
def test_automatic_memory_policy_and_candidate_segmentation(self) -> None:
policy = AutomaticMemoryPolicy(hidden_size=8)
output = policy(torch.randn(3, 8))
self.assertEqual(tuple(output.shape), (3,))
self.assertEqual(
split_memory_candidates("我叫林浩,我正在开发星火项目。"),
["我叫林浩,我正在开发星火项目。"],
)
self.assertTrue(looks_like_question("如果我选择 GPU,会发生什么?"))
self.assertFalse(looks_like_question("我住在上海,正在开发星火项目。"))
if __name__ == "__main__":
unittest.main()