Natural Memory NM2.1: 记忆路由器分叉、数据集缺陷修复与全轴评测证据

- 引入 MemoryRouterXL 与 v5/v6 流式多线程训练/编码管线
- 修复 prepare_memory_router_dataset 候选池重建缺陷(mega 家族 3568x 加速,输出逐字节相同)
- 修复 v5 被破坏的拒答与多跳标签(train 未知样本 319 -> 16319,multi_hop 平均正例 1.00 -> 2.00)
- 同存储预算下 V2-128 v6 逐轴 22/22 通过:Top-1 41.12% -> 94.62%,未知拒答 0.00% -> 100.00%
- 记录三条被实测推翻的显然优化(logits_to_keep=1 反而慢 55%、XL 容量未带来收益)
- 记忆手术跨架构可移植性 14/14,读写关闭时与原生模型逐位相同
This commit is contained in:
WpyQwq
2026-09-19 11:11:31 +08:00
commit 643e22ecb9
484 changed files with 306821 additions and 0 deletions
View File
+139
View File
@@ -0,0 +1,139 @@
"""Regression tests for the end-to-end answer scorer.
Both defects locked in here had produced headline numbers that were wrong enough
to send the project after the wrong bottleneck, so they get tests rather than
comments:
* whitespace sensitivity turned 19 correct multi-hop answers into failures
(multi-hop read 24.00% when it was 100.00%);
* keyword-only refusal detection missed 12 of 13 real refusals on the
unknown-attribute category (the axis read 0.00% when it was 52.00%).
"""
from __future__ import annotations
import unittest
from V2_dpskw.eval_scoring import is_refusal, score_case, squash
from V2_dpskw.eval_end_to_end_memory import evidence_write_order
def case(acceptable, answerable=True):
return {"acceptable": acceptable, "answerable": answerable}
class SquashTest(unittest.TestCase):
def test_removes_ascii_and_full_width_whitespace(self):
self.assertEqual(squash("值班人 -5259"), squash("值班人-5259"))
self.assertEqual(squash("值班人\u3000-5259"), squash("值班人-5259"))
self.assertEqual(squash("每周三 22:00"), squash("每周三 22:00"))
def test_is_case_insensitive(self):
self.assertEqual(squash("REF-9x6uye"), squash("ref-9X6UYE"))
def test_keeps_meaningful_punctuation(self):
# Punctuation inside anchors is meaningful and must survive.
self.assertNotEqual(squash("10.20.3.7:5432"), squash("1020375432"))
class WhitespaceInsensitiveMatchingTest(unittest.TestCase):
def test_spaced_identifier_still_scores_correct(self):
scored = score_case(case(["值班人-5259"]), "负责该项目的同事对应的值班人是 值班人 -5259。")
self.assertTrue(scored["correct"])
self.assertEqual(scored["matched"], ["值班人-5259"])
def test_wrong_identifier_is_still_wrong(self):
scored = score_case(case(["值班人-5259"]), "负责该项目的同事对应的值班人是 值班人-5266。")
self.assertFalse(scored["correct"])
def test_missing_value_is_wrong(self):
scored = score_case(case(["值班人-5259"]), "负责该项目的同事对应的值班人是 值班人。")
self.assertFalse(scored["correct"])
class RefusalDetectionTest(unittest.TestCase):
def test_recognises_the_wording_the_model_actually_uses(self):
for reply in (
"关于您的 weekly 会议安排,当前长期记忆中未包含相关信息,无法回答。",
"关于您这周按什么表轮,当前长期记忆中未包含该信息。已知信息如下:user;培训周期;78%",
"关于报警触发条件,当前记忆中没有相关记录。已知信息:user;保险到期;A-7719。",
"您的协议编号未知。",
"关于变更时段,当前长期记忆中未包含具体信息,证据不足,明确说不知道。",
):
with self.subTest(reply=reply):
self.assertTrue(is_refusal(reply))
def test_does_not_fire_on_asserted_values(self):
for reply in (
"您的培训周期是 32GB。",
"您的东西放在 63.2% 的仓库库位。",
"我固定每周三 22:00 做上线。",
):
with self.subTest(reply=reply):
self.assertFalse(is_refusal(reply))
def test_unanswerable_case_accepts_a_refusal_in_any_wording(self):
scored = score_case(case([], answerable=False), "当前长期记忆中未包含相关信息,无法回答。")
self.assertTrue(scored["correct"])
self.assertTrue(scored["abstained"])
def test_unanswerable_case_still_fails_when_a_value_is_asserted(self):
scored = score_case(case([], answerable=False), "您的培训周期是 32GB。")
self.assertFalse(scored["correct"])
def test_answerable_case_that_refuses_is_a_false_refusal(self):
scored = score_case(case(["32GB"]), "关于您上手所需的时间,当前长期记忆中未包含此信息,无法回答。")
self.assertFalse(scored["correct"])
self.assertTrue(scored["wrongly_abstained"])
def test_answerable_case_with_hedge_but_correct_value_is_not_false_refusal(self):
scored = score_case(
case(["凌晨 03:30"]),
"那杯偏爱 凌晨 03:30 的口味是未知,因为长期记忆中仅记录了时间(凌晨 03:30)和饮品名称。",
)
self.assertTrue(scored["correct"])
self.assertFalse(scored["wrongly_abstained"])
class EvidenceWriteOrderTest(unittest.TestCase):
"""The harness write order is a measurement decision, so it gets a test.
Two distinct failures lived here: answering facts written too early made
`update_conflict` unsolvable by construction, and reserving no slots for them
dropped the answer entirely on `noise_context` (72.00% -> 0.00%).
"""
def test_answering_facts_are_written_last(self):
case = {"positives": ["ANSWER"], "facts": ["d1", "d2", "ANSWER", "d3"]}
order = evidence_write_order(case, 6, answer_last=True)
self.assertEqual(order[-1], "ANSWER")
self.assertEqual(len(order), 4)
def test_answering_facts_survive_more_distractors_than_slots(self):
# noise_context: 8 distractors against 1 answering fact, 6 write slots.
case = {"positives": ["ANSWER"], "facts": ["d%d" % i for i in range(8)] + ["ANSWER"]}
order = evidence_write_order(case, 6, answer_last=True)
self.assertEqual(len(order), 6)
self.assertIn("ANSWER", order)
self.assertEqual(order[-1], "ANSWER")
self.assertEqual(order.count("ANSWER"), 1)
def test_all_answering_facts_survive(self):
case = {"positives": ["A1", "A2"], "facts": ["d1", "d2", "d3", "d4", "A1", "A2"]}
order = evidence_write_order(case, 6, answer_last=True)
self.assertIn("A1", order)
self.assertIn("A2", order)
self.assertEqual(order[-2:], ["A1", "A2"])
def test_answer_first_reproduces_the_old_behaviour(self):
case = {"positives": ["ANSWER"], "facts": ["d1", "d2", "ANSWER"]}
self.assertEqual(evidence_write_order(case, 6, answer_last=False), ["ANSWER", "d1", "d2"])
def test_never_exceeds_the_slot_budget(self):
case = {"positives": ["A"], "facts": ["d%d" % i for i in range(20)] + ["A"]}
self.assertLessEqual(len(evidence_write_order(case, 6)), 6)
self.assertLessEqual(len(evidence_write_order(case, 2)), 2)
if __name__ == "__main__":
unittest.main()
+172
View File
@@ -0,0 +1,172 @@
from __future__ import annotations
import unittest
import torch
from V2_dpskw.benchmark_dirty_real_corpus_4b import (
_failure_stage,
_is_refusal_strict,
)
from V2_dpskw.memory_os_v2 import MemoryRouterV2, PagedMemoryBankV2
from V2_dpskw.qwen_integration import (
_looks_like_grounded_memory_query,
format_memory_evidence,
)
class FailureDrivenV2Test(unittest.TestCase):
def setUp(self) -> None:
torch.manual_seed(19)
self.router = MemoryRouterV2(16, router_dim=8, num_heads=2, max_hops=3)
def test_grounding_guard_is_scoped_to_personal_and_project_facts(self) -> None:
self.assertTrue(_looks_like_grounded_memory_query("当前项目的模型叫什么"))
self.assertTrue(_looks_like_grounded_memory_query("README 是否记录了重启方式"))
self.assertFalse(_looks_like_grounded_memory_query("北京今天会下雨吗"))
def test_unknown_wording_is_a_hard_refusal_signal(self) -> None:
self.assertTrue(_is_refusal_strict("仓库中没有列出任何具体的部署区域。"))
self.assertTrue(_is_refusal_strict("这个字段没有出现在保存的记录里,我无法确认。"))
def test_structured_evidence_keeps_raw_fragment_first(self) -> None:
evidence = format_memory_evidence(
"用户给当前模型定名:Natural Memory v1。",
entity="user",
attribute="模型名称",
value="Natural Memory v1",
)
self.assertLess(evidence.index("事实"), evidence.index("已确认值"))
self.assertIn("已确认值:Natural Memory v1", evidence)
def test_failure_stage_classification_is_observable(self) -> None:
self.assertEqual(
_failure_stage(
{
"correct": False,
"answerable": True,
"retrieval_target_found": False,
"response": "不知道",
}
),
"routing_miss",
)
self.assertEqual(
_failure_stage(
{
"correct": False,
"answerable": True,
"retrieval_target_found": True,
"expected_anchors": ["H:\\Memory", "Natural Memory"],
"response": "项目在 H:\\Memory。",
}
),
"evidence_fusion",
)
self.assertEqual(
_failure_stage(
{
"correct": False,
"answerable": True,
"retrieval_target_found": True,
"expected_anchors": ["H:\\Memory"],
"response": "没有记录。",
}
),
"generation_control",
)
self.assertEqual(
_failure_stage(
{
"correct": False,
"answerable": False,
"retrieval_target_found": False,
"response": "我无法确认。",
}
),
"unknown_refusal_control",
)
def test_attribute_alias_routes_without_exact_source_wording(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
top_k_pages=2,
top_k_records=2,
max_hops=2,
)
record, _ = bank.write(
text="记忆主体优先放在进程内存,显存只保存热点记录。",
key=torch.randn(16),
entity="user",
attribute="记忆存储优先级",
value="DRAM",
confidence=0.98,
)
selected, _ = bank.query(
query_key=torch.randn(16),
query_text="容量不够时才降级,平时优先使用哪层内存?",
top_k_pages=2,
top_k_records=2,
)
self.assertIn(record.record_id, {item.record_id for item in selected})
def test_cjk_phrase_alias_reaches_operational_record(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
top_k_pages=2,
top_k_records=2,
)
record, _ = bank.write(
text="当前 token 不会对一百万个 slot 做全量注意力。",
key=torch.randn(16),
semantic_key=torch.randn(16),
entity="README_NATURAL_MEMORY_V2.md",
attribute="operational:1",
value="Top-K;全量注意力",
confidence=0.98,
)
selected, _ = bank.query(
query_key=torch.randn(16),
query_text="百万 slot 会不会进行全量注意力?",
top_k_pages=2,
top_k_records=2,
)
self.assertIn(record.record_id, {item.record_id for item in selected})
def test_ambiguous_symbol_fans_out_but_stays_bounded(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
top_k_pages=8,
top_k_records=1,
max_pages=64,
)
expected = set()
for index in range(6):
record, _ = bank.write(
text=f"第 {index} 个文件定义了 run 函数。",
key=torch.randn(16),
entity=f"module_{index}.py",
attribute="symbol:run",
value="run",
confidence=0.9,
)
expected.add(record.record_id)
selected, _ = bank.query(
query_key=torch.randn(16),
query_text="run 函数分别在哪些文件中定义?",
top_k_pages=8,
top_k_records=1,
)
selected_ids = {item.record_id for item in selected}
self.assertEqual(selected_ids, expected)
self.assertLessEqual(len(selected), 8)
if __name__ == "__main__":
unittest.main()
+473
View File
@@ -0,0 +1,473 @@
"""Retraction must be turn-scoped, not just key-scoped.
The bug this covers was measured end-to-end, not reasoned about. A user said
"remember: my emergency contact is Wang, extension 7781", which the agent stored as two
records (contact and extension). The user then said "forget my emergency contact
information". The router's auto-forget retracted only the record whose key matched --
the contact -- leaving the extension active, and on the next turn the model read that
leftover record back and answered "your emergency contact extension is 7781": a fact the
user had just revoked.
Records now carry the ``origin`` of the user turn they came from, and retraction can
retire a whole origin group, so the unit of forgetting is the unit of telling.
"""
from __future__ import annotations
import sqlite3
import unittest
from dataclasses import asdict
from pathlib import Path
from tempfile import TemporaryDirectory
import torch
from V2_dpskw.memory_os_v2 import (
MemoryOSV2,
MemoryRecordV2,
MemoryRouterV2,
PagedMemoryBankV2,
STATUS_ACTIVE,
STATUS_RETRACTED,
memory_record_to_dict,
)
from V2_dpskw.qwen_integration import (
compose_corrected_evidence,
infer_memory_metadata,
memory_fact_sentence,
memory_origin,
)
from V2_dpskw.tiered_memory_store_v2 import TieredMemoryStoreV2, _pack_tensor
TURN = "记一下:我的紧急联系人是 王工,电话分机 7781。"
OTHER_TURN = "记一下:我的常用编辑器是 VSCode。"
def _bank(key_dim: int = 16, **kwargs) -> PagedMemoryBankV2:
return PagedMemoryBankV2(
key_dim,
router=MemoryRouterV2(key_dim, router_dim=8, num_heads=2, max_hops=3),
page_capacity=8,
max_pages=512,
hot_pages=4,
top_k_pages=4,
top_k_records=8,
max_hops=3,
coarse_index_bits=8,
**kwargs,
)
class MemoryOriginForgetTest(unittest.TestCase):
def setUp(self) -> None:
torch.manual_seed(11)
self.bank = _bank()
def _write_turn(self, origin: str, **fields):
return self.bank.write(
key=torch.randn(16),
origin=origin,
confidence=0.95,
**fields,
)[0]
def test_forget_retires_every_record_from_the_turn(self) -> None:
"""The measured leak: the sibling record must not survive the forget."""
origin = memory_origin(TURN)
contact = self._write_turn(
origin, text="user 的 紧急联系人 是 王工。", entity="user",
attribute="紧急联系人", value="王工",
)
extension = self._write_turn(
origin, text="user 的 电话分机 是 7781。", entity="user",
attribute="电话分机", value="7781",
)
self.assertEqual(contact.status, STATUS_ACTIVE)
self.assertEqual(extension.status, STATUS_ACTIVE)
# Forgetting "the emergency contact information" addresses the contact record.
retracted = self.bank.retract_origin(contact.origin)
self.assertIn(contact.record_id, retracted)
self.assertIn(extension.record_id, retracted)
self.assertEqual(contact.status, STATUS_RETRACTED)
self.assertEqual(extension.status, STATUS_RETRACTED)
self.assertEqual(
[r for r in self.bank.records.values() if r.status == STATUS_ACTIVE], []
)
def test_retract_origin_spares_other_turns(self) -> None:
target = self._write_turn(
memory_origin(TURN), text="user 的 紧急联系人 是 王工。",
entity="user", attribute="紧急联系人", value="王工",
)
bystander = self._write_turn(
memory_origin(OTHER_TURN), text="user 的 常用编辑器 是 VSCode。",
entity="user", attribute="常用编辑器", value="VSCode",
)
self.bank.retract_origin(target.origin)
self.assertEqual(target.status, STATUS_RETRACTED)
self.assertEqual(bystander.status, STATUS_ACTIVE)
def test_empty_origin_never_matches_anything(self) -> None:
"""Records predating this field carry no origin and must be untouchable."""
# A blank turn must not hash to a *shared* non-empty origin, or every
# text-less record would form one group and forgetting one would take all.
self.assertEqual(memory_origin(""), "")
self.assertEqual(memory_origin(" "), "")
self.assertEqual(memory_origin(None), "")
legacy = self._write_turn("", text="旧记录", entity="user", attribute="旧", value="x")
blank = self._write_turn(
memory_origin(""), text="", entity="user", attribute="空", value="x"
)
self.assertEqual(self.bank.retract_origin(""), [])
self.assertEqual(self.bank.retract_origin(" "), [])
self.assertEqual(legacy.status, STATUS_ACTIVE)
self.assertEqual(blank.status, STATUS_ACTIVE)
# A real origin must not sweep up the origin-less record either.
self.assertEqual(self.bank.retract_origin(memory_origin(TURN)), [])
self.assertEqual(legacy.status, STATUS_ACTIVE)
self.assertEqual(blank.status, STATUS_ACTIVE)
def test_edit_record_inherits_origin(self) -> None:
origin = memory_origin(TURN)
original = self._write_turn(
origin, text="user 的 紧急联系人 是 王工。", entity="user",
attribute="紧急联系人", value="王工",
)
edited = self.bank.edit_record(original.record_id, value="李工")
self.assertEqual(edited.origin, origin)
self.assertEqual(edited.value, "李工")
# The group still retracts as a unit after versioning.
self.assertEqual(len(self.bank.retract_origin(origin)), 1)
self.assertEqual(edited.status, STATUS_RETRACTED)
def test_record_view_exposes_origin(self) -> None:
record = self._write_turn(
memory_origin(TURN), text="user 的 紧急联系人 是 王工。", entity="user",
attribute="紧急联系人", value="王工",
)
self.assertEqual(memory_record_to_dict(record)["origin"], memory_origin(TURN))
def test_payload_round_trip_and_legacy_payloads(self) -> None:
"""``asdict`` is what export serialises; ``from_payload`` splats the dict back."""
record = self._write_turn(
memory_origin(TURN), text="user 的 紧急联系人 是 王工。", entity="user",
attribute="紧急联系人", value="王工",
)
payload = asdict(record)
self.assertEqual(MemoryRecordV2(**payload).origin, memory_origin(TURN))
# A payload written before the field existed must still load.
payload.pop("origin")
self.assertEqual(MemoryRecordV2(**payload).origin, "")
def test_origin_survives_the_tier_store(self) -> None:
"""The SQLite store keeps its own column list, so this is the real risk."""
origin = memory_origin(TURN)
with TemporaryDirectory() as tmp:
path = Path(tmp) / "store.db"
store = TieredMemoryStoreV2(path, key_dim=16, page_capacity=8)
bank = _bank(tier_store=store)
first = bank.write(
text="user 的 紧急联系人 是 王工。", key=torch.randn(16), origin=origin,
entity="user", attribute="紧急联系人", value="王工", confidence=0.95,
)[0]
second = bank.write(
text="user 的 电话分机 是 7781。", key=torch.randn(16), origin=origin,
entity="user", attribute="电话分机", value="7781", confidence=0.95,
)[0]
store.close()
# A fresh bank over the same file must see the origin it never had in RAM.
reopened_store = TieredMemoryStoreV2(path, key_dim=16, page_capacity=8)
reloaded = _bank(tier_store=reopened_store)
reloaded._hydrate_records([first.record_id, second.record_id])
for record_id in (first.record_id, second.record_id):
self.assertEqual(
reloaded.records[record_id].origin, origin,
"origin was lost on the way through the store",
)
retracted = reloaded.retract_origin(origin)
self.assertEqual(len(retracted), 2)
reopened_store.close()
def test_store_migrates_a_database_written_without_origin(self) -> None:
"""An existing store has no ``origin`` column; opening it must widen it."""
with TemporaryDirectory() as tmp:
path = Path(tmp) / "old.db"
# Reproduce a pre-change store: same schema, no origin column.
connection = sqlite3.connect(str(path))
connection.executescript(
"""
CREATE TABLE store_meta (key TEXT PRIMARY KEY, value TEXT NOT NULL);
CREATE TABLE records (
record_id TEXT PRIMARY KEY, page_id TEXT NOT NULL, text TEXT NOT NULL,
key BLOB NOT NULL, summary BLOB NOT NULL, memory_type TEXT NOT NULL,
entity TEXT NOT NULL, attribute TEXT NOT NULL, value TEXT NOT NULL,
timestamp INTEGER NOT NULL, importance REAL NOT NULL,
confidence REAL NOT NULL, source TEXT NOT NULL, status TEXT NOT NULL,
version INTEGER NOT NULL, supersedes TEXT NOT NULL,
related_ids TEXT NOT NULL, evidence TEXT NOT NULL,
slot_index INTEGER NOT NULL, token_ids BLOB, token_mask BLOB,
access_count INTEGER NOT NULL, last_access INTEGER NOT NULL
);
"""
)
connection.execute(
"INSERT INTO records(record_id, page_id, text, key, summary, memory_type,"
" entity, attribute, value, timestamp, importance, confidence, source,"
" status, version, supersedes, related_ids, evidence, slot_index,"
" token_ids, token_mask, access_count, last_access) VALUES"
"('legacy', 'page_00000001', '旧记录', ?, ?, 'fact', 'user', '旧',"
" 'x', 1, 0.5, 0.5, 'user', 'active', 0, '', '[]', '[]', -1, NULL, NULL, 0, 1)",
(
_pack_tensor(torch.zeros(16), dtype="float32"),
_pack_tensor(torch.zeros(16), dtype="float32"),
),
)
connection.commit()
connection.close()
store = TieredMemoryStoreV2(path, key_dim=16, page_capacity=8)
columns = {str(row[1]) for row in store.connection.execute("PRAGMA table_info(records)")}
self.assertIn("origin", columns)
# The row written before the column existed reads back as origin-less, and
# the widened table still accepts ordinary writes.
self.assertEqual(store.load_records(["legacy"])[0]["origin"], "")
bank = _bank(tier_store=store)
record = bank.write(
text="新记录", key=torch.randn(16), origin=memory_origin(TURN),
entity="user", attribute="新", value="y", confidence=0.95,
)[0]
self.assertEqual(store.load_records([record.record_id])[0]["origin"], memory_origin(TURN))
store.close()
def test_structured_write_absorbs_the_turns_unidentified_record(self) -> None:
"""The auto layer writes before the model generates, so it cannot know the turn's
identity. A structured write for the same turn must retire that raw record."""
origin = memory_origin(TURN)
# The automatic layer's record: whole sentence, no entity/attribute inferred.
raw = self._write_turn(
origin, text="【长期记忆证据】 事实:我叫Wpy", entity="", attribute="", value="",
)
self.assertEqual(raw.conflict_key(), "")
self.assertEqual(raw.status, STATUS_ACTIVE)
# The agent's structured write for the same turn.
structured = self._write_turn(
origin, text="我的名字是 Wpy", entity="user", attribute="名字", value="Wpy",
)
self.assertEqual(structured.status, STATUS_ACTIVE)
self.assertEqual(raw.status, "superseded")
active = [r for r in self.bank.records.values() if r.status == STATUS_ACTIVE]
self.assertEqual([r.record_id for r in active], [structured.record_id])
def test_absorption_is_confined_to_the_same_turn(self) -> None:
"""This runs on every structured write, so it must not reach other turns."""
raw_other = self._write_turn(
memory_origin(OTHER_TURN), text="【长期记忆证据】 事实:我住在上海",
entity="", attribute="", value="",
)
structured = self._write_turn(
memory_origin(TURN), text="我的名字是 Wpy", entity="user", attribute="名字", value="Wpy",
)
self.assertEqual(structured.status, STATUS_ACTIVE)
self.assertEqual(raw_other.status, STATUS_ACTIVE)
def test_absorption_spares_identified_siblings(self) -> None:
"""One turn can carry two facts; the second structured write must not eat the first."""
origin = memory_origin(TURN)
contact = self._write_turn(
origin, text="紧急联系人是 王工", entity="user", attribute="紧急联系人", value="王工",
)
extension = self._write_turn(
origin, text="电话分机是 7781", entity="user", attribute="电话分机", value="7781",
)
self.assertEqual(contact.status, STATUS_ACTIVE)
self.assertEqual(extension.status, STATUS_ACTIVE)
def test_unidentified_write_absorbs_nothing(self) -> None:
"""Only a record that *has* an identity may absorb, or the raw records would
supersede each other and the turn's fact would be lost."""
origin = memory_origin(TURN)
first = self._write_turn(origin, text="事实:我叫Wpy", entity="", attribute="", value="")
second = self._write_turn(origin, text="事实:我是开发者", entity="", attribute="", value="")
self.assertEqual(first.status, STATUS_ACTIVE)
self.assertEqual(second.status, STATUS_ACTIVE)
def test_absorption_then_correction_leaves_one_active_record(self) -> None:
"""The end-to-end shape of the bug: after absorption a correction has a key to
version, so the stale value cannot survive next to the new one."""
origin = memory_origin(TURN)
self._write_turn(origin, text="事实:我叫Wpy", entity="", attribute="", value="")
structured = self._write_turn(
origin, text="我的名字是 Wpy", entity="user", attribute="名字", value="Wpy",
)
corrected = self.bank.edit_record(structured.record_id, value="王五")
active = [r for r in self.bank.records.values() if r.status == STATUS_ACTIVE]
self.assertEqual([r.record_id for r in active], [corrected.record_id])
self.assertEqual(corrected.value, "王五")
self.assertEqual(corrected.version, 1)
# The stale value is gone from the active set: the raw record that carried it was
# absorbed, so nothing active still answers "Wpy".
self.assertEqual([r.value for r in active], ["王五"])
# Note for future readers: ``edit_record(value=...)`` keeps the old ``text`` on
# purpose, so an edit that only changes the value leaves a stale sentence behind --
# and the memory layer quotes ``text`` back to the model as the 事实 line. Callers
# that need the sentence to follow the value pass ``text`` as well; the agent's
# nm2_write tool does exactly that, which is why the TUI shows the corrected sentence.
self.assertEqual(corrected.text, "我的名字是 Wpy")
def test_memory_os_v2_forwards_retract_origin(self) -> None:
"""The model-facing facade is what the integration actually calls."""
os_v2 = MemoryOSV2(16, router=self.bank.router, bank=self.bank)
origin = memory_origin(TURN)
for attribute, value in (("紧急联系人", "王工"), ("电话分机", "7781")):
os_v2.write(
text=f"user 的 {attribute} 是 {value}。",
key=torch.randn(16),
entity="user",
attribute=attribute,
value=value,
confidence=0.95,
origin=origin,
)
retracted = os_v2.retract_origin(origin)
self.assertEqual(len(retracted), 2)
self.assertEqual(os_v2.stats()["active_records"], 0)
class CorrectedEvidenceTest(unittest.TestCase):
"""A correction must reach the thing the model actually reads.
The model reads the record's stored evidence card, not its structured fields. A
correction that supplied only a new value therefore used to report success, bump the
version and supersede the old record while every injected card still said
``已确认值:<old>`` -- and the model kept answering the old value.
"""
CARD = (
"【长期记忆证据】\n"
"事实:我的名字是 Wpy\n"
"实体:user\n属性:名字\n已确认值:Wpy\n"
"可直接复述的关键短语:user;名字;Wpy\n"
"回答要求:第一句先直接回答问题,并原样复述上面的关键短语。"
)
def test_value_correction_removes_the_old_value_everywhere(self) -> None:
rebuilt = compose_corrected_evidence(
{"entity": "user", "attribute": "名字", "value": "Wpy", "text": self.CARD},
value="王五",
)
self.assertIsNotNone(rebuilt)
assert rebuilt is not None
self.assertNotIn("Wpy", rebuilt)
self.assertIn("事实:我的名字是 王五", rebuilt)
self.assertIn("已确认值:王五", rebuilt)
self.assertIn("关键短语:user;名字;王五", rebuilt)
def test_the_sentence_is_reused_and_natural(self) -> None:
"""Substituting in place keeps the original phrasing rather than emitting labels."""
rebuilt = compose_corrected_evidence(
{"entity": "user", "attribute": "名字", "value": "Wpy", "text": self.CARD},
value="王五",
)
assert rebuilt is not None
self.assertIn("事实:我的名字是 王五", rebuilt)
def test_value_absent_from_the_sentence_falls_back(self) -> None:
rebuilt = compose_corrected_evidence(
{"entity": "user", "attribute": "分机", "value": "7781", "text": "记一下:分机号码是它。"},
value="9999",
)
assert rebuilt is not None
self.assertNotIn("7781", rebuilt)
self.assertIn("9999", rebuilt)
def test_nothing_to_rebuild_returns_none(self) -> None:
"""No field supplied means the caller is not correcting a fact -- leave it alone."""
self.assertIsNone(compose_corrected_evidence(
{"entity": "user", "attribute": "名字", "value": "Wpy", "text": self.CARD},
))
# A record with no structured identity cannot be rendered better than it already is.
self.assertIsNone(compose_corrected_evidence({"entity": "", "attribute": "", "value": "", "text": "x"}, value="y"))
self.assertIsNone(compose_corrected_evidence({"entity": "user", "attribute": "名字", "value": "Wpy", "text": "x"}, value=""))
def test_plain_sentence_records_are_handled(self) -> None:
"""Records written by the agent hold a plain sentence, not a card."""
self.assertEqual(memory_fact_sentence("记一下:我的名字是 王五。"), "记一下:我的名字是 王五。")
rebuilt = compose_corrected_evidence(
{"entity": "user", "attribute": "名字", "value": "Wpy", "text": "我的名字是 Wpy。"},
value="王五",
)
assert rebuilt is not None
self.assertIn("事实:我的名字是 王五。", rebuilt)
self.assertNotIn("Wpy", rebuilt.replace("回答要求", ""))
class AttributeExtractionTest(unittest.TestCase):
"""A drive-letter colon is not a key/value separator.
Measured: 「记住:我的评测报告放在 E:\\deepseek\\artifacts。」 states a place with 放在
rather than 是, so the parser fell through to the bare colon in its operator list and split
on the drive letter, filing the fact under attribute 「评测报告放在 E」 with value
「\\deepseek\\artifacts」. The same fact written by the agent used 「评测报告路径」, so the
two keys could never match, could never version each other, and the stale path stayed
active and ranked above the corrected one.
"""
def test_paths_are_not_split_on_the_drive_colon(self) -> None:
for sentence in (
"记住:我的评测报告放在 E:\\deepseek\\artifacts。",
"改了:我的评测报告以后放到 H:\\Memory\\agent_lab\\runs。",
):
meta = infer_memory_metadata(sentence)
self.assertNotIn("E", str(meta.get("attribute")), sentence)
self.assertNotIn("H", str(meta.get("attribute")), sentence)
# No attribute is better than a wrong one: it stays an unstructured episode rather
# than claiming a conflict key that nothing else can match.
self.assertIsNone(meta.get("attribute"), sentence)
def test_a_colon_after_cjk_is_still_a_separator(self) -> None:
meta = infer_memory_metadata("我的密钥:abc")
self.assertEqual(meta.get("attribute"), "密钥")
self.assertEqual(meta.get("value"), "abc")
def test_the_copula_still_wins_for_paths(self) -> None:
meta = infer_memory_metadata("我的评测报告目录是 E:\\deepseek\\artifacts")
self.assertEqual(meta.get("attribute"), "评测报告目录")
self.assertEqual(meta.get("value"), "E:\\deepseek\\artifacts")
def test_colons_inside_values_survive(self) -> None:
self.assertEqual(infer_memory_metadata("我的时区是 UTC+8").get("value"), "UTC+8")
self.assertEqual(infer_memory_metadata("我的密钥是 abc:def").get("value"), "abc:def")
if __name__ == "__main__":
unittest.main()
+594
View File
@@ -0,0 +1,594 @@
from __future__ import annotations
import unittest
from tempfile import TemporaryDirectory
import torch
from V2_dpskw.memory_os_v2 import (
KVBudgetManagerV2,
MemoryOSV2,
MemoryRouterV2,
PagedMemoryBankV2,
STATUS_ACTIVE,
STATUS_QUARANTINED,
STATUS_SUPERSEDED,
)
from V2_dpskw.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_independent_fragments_keep_semantic_keys_across_export(self) -> None:
first, _ = self.bank.write(
text="项目负责人是成员A",
key=torch.randn(16),
semantic_key=torch.randn(16),
slot_index=-1,
token_ids=torch.tensor([1, 2, 3]),
)
second, _ = self.bank.write(
text="成员A的工作代号是H7",
key=torch.randn(16),
semantic_key=torch.randn(16),
slot_index=-1,
token_ids=torch.tensor([4, 5, 6]),
)
self.assertEqual(self.bank.stats()["active_records"], 2)
payload = self.bank.export_payload()
restored = PagedMemoryBankV2.from_payload(payload, router=self.router)
self.assertEqual(restored.stats()["active_records"], 2)
self.assertIsNotNone(restored.records[first.record_id].semantic_key)
self.assertIsNotNone(restored.records[second.record_id].semantic_key)
def test_bounded_record_reranker_runs_inside_selected_pages(self) -> None:
first, _ = self.bank.write(
text="候选一",
key=torch.tensor([1.0] + [0.0] * 15),
semantic_key=torch.ones(16),
)
second, _ = self.bank.write(
text="候选二",
key=torch.tensor([1.0] + [0.0] * 15),
semantic_key=torch.ones(16) * 2,
)
def scorer(query, candidates):
# The callback receives only the two records in the selected page.
return torch.tensor([0.1, 0.9], device=query.device)
self.bank.record_scorer = scorer
records, _ = self.bank.query(
query_key=torch.tensor([1.0] + [0.0] * 15),
query_text="候选",
top_k_pages=1,
top_k_records=1,
)
self.assertEqual(records[0].record_id, second.record_id)
def test_explicit_entity_address_beats_semantic_collision(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
hot_pages=0,
top_k_pages=1,
top_k_records=1,
coarse_index_bits=8,
)
distractor, _ = bank.write(
text="评估用户00056的档案代号为x",
key=torch.ones(16),
entity="评估用户00056",
attribute="档案代号",
value="x",
)
target, _ = bank.write(
text="评估用户00062的档案代号为w",
key=torch.ones(16),
entity="评估用户00062",
attribute="档案代号",
value="w",
)
records, _ = bank.query(
query_key=torch.ones(16),
query_text="只查询评估用户00062的档案代号",
top_k_pages=1,
top_k_records=1,
)
self.assertEqual(records[0].record_id, target.record_id)
self.assertNotEqual(records[0].record_id, distractor.record_id)
def test_symbol_address_returns_bounded_ambiguous_candidates(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
hot_pages=0,
top_k_pages=4,
top_k_records=1,
coarse_index_bits=8,
)
records = []
for path in ("first.py", "second.py", "third.py"):
record, _ = bank.write(
text=f"def run in {path}",
key=torch.ones(16),
semantic_key=torch.ones(16),
entity=path,
attribute="symbol:run",
value=path,
token_ids=torch.tensor([1, 2, 3]),
)
records.append(record)
selected, _ = bank.query(
query_key=torch.ones(16),
query_text="帮我定位项目里的 run 函数",
top_k_pages=4,
top_k_records=1,
)
self.assertEqual(
{record.record_id for record in selected},
{record.record_id for record in records},
)
def test_explicit_file_address_filters_same_named_symbol(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
hot_pages=0,
top_k_pages=8,
top_k_records=1,
coarse_index_bits=8,
)
target, _ = bank.write(
text="def main in target.py",
key=torch.ones(16),
semantic_key=torch.ones(16),
entity="target.py",
attribute="symbol:main",
value="target.py",
token_ids=torch.tensor([1, 2, 3]),
)
distractor, _ = bank.write(
text="def main in other.py",
key=torch.ones(16),
semantic_key=torch.ones(16),
entity="other.py",
attribute="symbol:main",
value="other.py",
token_ids=torch.tensor([1, 2, 3]),
)
selected, _ = bank.query(
query_key=torch.ones(16),
query_text="请定位 target.py 里的 main 函数",
top_k_pages=8,
top_k_records=1,
)
self.assertEqual([record.record_id for record in selected], [target.record_id])
self.assertNotIn(distractor.record_id, {record.record_id for record in selected})
def test_identifier_subtokens_reach_operational_evidence(self) -> None:
target, _ = self.bank.write(
text="代码定义 DEFAULT_MEMORY_RESET_TOKEN,用于清空记忆",
key=torch.ones(16),
semantic_key=torch.ones(16),
entity="qwen_integration.py",
attribute="operational:reset",
value="qwen_integration.py",
token_ids=torch.tensor([1, 2, 3]),
)
selected, _ = self.bank.query(
query_key=torch.ones(16),
query_text="模型代码里的默认 reset token 是什么",
top_k_pages=2,
top_k_records=1,
)
self.assertTrue(selected)
self.assertEqual(selected[0].record_id, target.record_id)
def test_exact_entity_attribute_filters_injected_distractor(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=8,
hot_pages=0,
top_k_pages=1,
top_k_records=2,
coarse_index_bits=8,
)
target, _ = bank.write(
text="评估用户00119的常用语言为X",
key=torch.ones(16),
entity="评估用户00119",
attribute="常用语言",
value="X",
)
distractor, _ = bank.write(
text="评估用户00009的常用语言为T",
key=torch.ones(16),
entity="评估用户00009",
attribute="常用语言",
value="T",
)
records, _ = bank.query(
query_key=torch.ones(16),
query_text="请查询评估用户00119的常用语言",
top_k_pages=1,
top_k_records=2,
)
self.assertEqual([record.record_id for record in records], [target.record_id])
self.assertNotIn(distractor.record_id, [record.record_id for record in records])
def test_distinctive_address_reaches_target_beyond_coarse_page_limit(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
max_pages=512,
hot_pages=0,
top_k_pages=1,
top_k_records=1,
coarse_index_bits=8,
)
target = None
for index in range(180):
record, _ = bank.write(
text=f"代码文件 file_{index}.py 的函数 fn_{index}",
key=torch.ones(16),
entity=f"file_{index}.py",
attribute=f"symbol:fn_{index}",
value=f"file_{index}.py",
)
if index == 179:
target = record
assert target is not None
records, _ = bank.query(
query_key=torch.ones(16),
query_text="请查找 file_179.py 中的 fn_179 定义",
top_k_pages=1,
top_k_records=1,
)
self.assertTrue(records)
self.assertEqual(records[0].record_id, target.record_id)
def test_numeric_suffix_does_not_cross_entity_namespace(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=1,
hot_pages=0,
top_k_pages=1,
top_k_records=2,
coarse_index_bits=8,
)
training, _ = bank.write(
text="训练用户00006的档案代号为F",
key=torch.ones(16),
entity="训练用户00006",
attribute="档案代号",
value="F",
)
evaluation, _ = bank.write(
text="评估用户00006的档案代号为L",
key=torch.ones(16),
entity="评估用户00006",
attribute="档案代号",
value="L",
)
records, _ = bank.query(
query_key=torch.ones(16),
query_text="请读取评估用户00006的档案代号",
top_k_pages=1,
top_k_records=2,
)
self.assertIn(evaluation.record_id, [record.record_id for record in records])
self.assertNotIn(training.record_id, [record.record_id for record in records])
def test_explicit_address_is_marked_for_fast_route(self) -> None:
bank = PagedMemoryBankV2(
16,
router=self.router,
page_capacity=2,
max_pages=16,
coarse_index_bits=8,
)
bank.write(
text="代码文件 symbol:parse_args 的状态为enabled",
key=torch.randn(16),
entity="symbol:parse_args",
attribute="状态",
value="enabled",
)
self.assertTrue(bank.has_explicit_address("symbol:parse_args 当前状态是什么?"))
self.assertFalse(bank.has_explicit_address("当前状态是什么?"))
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 V2_dpskw.model import DynamicMemoryConfig, DynamicMemoryLM
from V2_dpskw.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()
+33
View File
@@ -0,0 +1,33 @@
import unittest
from V2_dpskw.qwen_integration import infer_memory_metadata
class NaturalMemorySafetyTest(unittest.TestCase):
def test_explicit_fact_gets_structured_version_metadata(self) -> None:
item = infer_memory_metadata("请记住:我的常用时区是Asia/Shanghai。")
self.assertEqual(item["kind"], "fact")
self.assertEqual(item["entity"], "user")
self.assertEqual(item["attribute"], "常用时区")
self.assertEqual(item["value"], "Asia/Shanghai")
self.assertTrue(item["should_write"])
def test_correction_is_structured_but_not_collapsed_into_a_question(self) -> None:
item = infer_memory_metadata("更正一下:我的工作地点改为深圳,旧值不再有效。")
self.assertEqual(item["kind"], "correction")
self.assertEqual(item["attribute"], "工作地点")
self.assertEqual(item["value"], "深圳")
def test_explicit_forget_never_becomes_a_new_record(self) -> None:
item = infer_memory_metadata("请删除关于我的水果偏好的记忆。")
self.assertEqual(item["kind"], "forget")
self.assertFalse(item["should_write"])
self.assertEqual(item["attribute"], "水果偏好")
def test_question_and_hypothetical_are_not_facts(self) -> None:
self.assertEqual(infer_memory_metadata("我的常住城市是什么?")["kind"], "query")
self.assertEqual(infer_memory_metadata("如果我的常住城市改成成都,会怎样?")["kind"], "query")
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 V2_dpskw.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()
+128
View File
@@ -0,0 +1,128 @@
from __future__ import annotations
import json
import unittest
from pathlib import Path
import torch
from V2_dpskw.memory_os_v2 import MemoryRouterV2
from V2_dpskw.prepare_memory_router_dataset import Corpus, _add_generic_qa_row, _add_normalized_policy_file, _source_candidates
from V2_dpskw.train_memory_router_large import _episode_tensors, _router_loss
class RouterTrainingContractTests(unittest.TestCase):
def test_policy_builder_keeps_latest_fact_and_excludes_query_candidates(self) -> None:
corpus = Corpus()
rows = [
(1, {"id": "g:fact", "group_id": "g", "text": "我的项目是旧项目。", "kind": "fact", "write_label": 1.0, "attribute": "项目", "subject": "u"}),
(2, {"id": "g:replacement", "group_id": "g", "text": "更正:我的项目改为新项目。", "kind": "replacement", "write_label": 1.0, "attribute": "项目", "subject": "u"}),
(3, {"id": "g:query", "group_id": "g", "text": "我的项目是什么?", "kind": "query", "write_label": 0.0, "attribute": "项目", "subject": "u"}),
(4, {"id": "g:forget", "group_id": "g", "text": "忘掉我的项目。", "kind": "forget", "write_label": 0.0, "attribute": "项目", "subject": "u"}),
]
_add_normalized_policy_file(corpus, "train", "local.jsonl", rows)
self.assertEqual(len(corpus.episodes["train"]), 1)
episode = corpus.episodes["train"][0]
row = _source_candidates(
"train",
episode,
corpus.items["train"],
candidate_count=4,
seed=7,
eligible_items=[item for item in corpus.items["train"].values() if item.kind not in {"query", "forget"}],
)
self.assertIsNotNone(row)
assert row is not None
candidates = row["candidates"]
self.assertTrue(all(item["kind"] not in {"query", "forget"} for item in candidates))
positive = candidates[row["positive_index"]]
self.assertIn("新项目", positive["text"])
def test_episode_tensors_support_variable_candidates_and_unknowns(self) -> None:
episodes = [
{
"query": "q1",
"candidates": [{"text": "a"}, {"text": "b"}],
"positive_indices": [1],
"need_memory": 1.0,
"hop": 1,
"family": "qa",
},
{
"query": "q2",
"candidates": [{"text": "c"}],
"positive_indices": [],
"need_memory": 0.0,
"hop": 0,
"family": "qa",
},
]
from V2_dpskw.train_memory_router_large import _text_key
lookup = {_text_key(text): index for index, text in enumerate(("q1", "a", "b", "q2", "c"))}
data = _episode_tensors(episodes, lookup)
self.assertEqual(tuple(data["candidate_indices"].shape), (2, 2))
self.assertEqual(data["candidate_mask"].tolist(), [[True, True], [True, False]])
self.assertEqual(data["positive_mask"].tolist(), [[False, True], [False, False]])
def test_router_loss_is_finite_for_multi_positive_and_abstention_batch(self) -> None:
torch.manual_seed(4)
router = MemoryRouterV2(16, router_dim=8, num_heads=2, max_hops=3)
batch = {
"query": torch.randn(2, 16),
"candidates": torch.randn(2, 3, 16),
"candidate_mask": torch.tensor([[True, True, True], [True, True, False]]),
"positive_mask": torch.tensor([[True, True, False], [False, False, False]]),
"need": torch.tensor([1.0, 0.0]),
"hops": torch.tensor([2, 0]),
}
args = type(
"Args",
(),
{
"margin": 0.10,
"need_loss_weight": 0.75,
"hop_loss_weight": 0.35,
"margin_loss_weight": 0.25,
},
)()
loss, parts = _router_loss(router, batch, args)
self.assertTrue(torch.isfinite(loss).item())
self.assertTrue(all(torch.isfinite(torch.tensor(value)).item() for value in parts.values()))
def test_codesearchnet_style_documentation_becomes_query(self) -> None:
corpus = Corpus()
_add_generic_qa_row(
corpus,
"train",
"code_search_net|train|python|train",
"hf:code_search_net",
1,
{
"id": "code-1",
"func_documentation_string": "load the memory page by id",
"whole_func_string": "def load_page(page_id): return pages[page_id]",
"repository_name": "example/repo",
},
"code",
)
self.assertEqual(len(corpus.episodes["train"]), 1)
self.assertEqual(corpus.episodes["train"][0].query, "load the memory page by id")
self.assertEqual(len(corpus.episodes["train"][0].positive_ids), 1)
def test_generated_manifest_and_eval_hash_are_consistent(self) -> None:
root = Path(__file__).resolve().parents[1] / "data" / "router_training"
manifest_path = root / "manifest.json"
eval_path = root / "eval.jsonl"
hash_path = root / "eval.sha256"
if not (manifest_path.exists() and eval_path.exists() and hash_path.exists()):
self.skipTest("router_training data has not been generated")
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
actual = __import__("hashlib").sha256(eval_path.read_bytes()).hexdigest()
self.assertEqual(actual, manifest["files"]["eval"]["sha256"])
self.assertEqual(actual, hash_path.read_text(encoding="ascii").strip())
self.assertEqual(manifest["leakage_check"]["group_overlap"], 0)
if __name__ == "__main__":
unittest.main()
+149
View File
@@ -0,0 +1,149 @@
"""Contract tests for the new MemoryRouterXL architecture.
The XL router must satisfy two things at once:
* be a genuinely larger, separate model (not a tweak of MemoryRouterV2), and
* expose the exact runtime contract ``PagedMemoryBankV2`` already drives, so a
trained XL router can be dropped into the existing memory OS.
These tests run on tiny tensors only; no Qwen model is loaded.
"""
from __future__ import annotations
import unittest
from pathlib import Path
from tempfile import TemporaryDirectory
import torch
from V2_dpskw.memory_os_v2 import MemoryRouterV2, PagedMemoryBankV2
from V2_dpskw.router_xl import ARCH_NAME, MemoryRouterXL, load_router_xl
def _tiny_xl(**overrides: object) -> MemoryRouterXL:
kwargs: dict[str, object] = dict(
router_dim=8,
num_heads=2,
max_hops=3,
encoder_layers=2,
encoder_hidden=8,
pair_blocks=1,
pair_hidden=8,
policy_layers=2,
policy_hidden=8,
)
kwargs.update(overrides)
torch.manual_seed(11)
return MemoryRouterXL(16, **kwargs)
class RouterXLContractTest(unittest.TestCase):
def test_forward_matches_the_runtime_contract(self) -> None:
router = _tiny_xl()
query = torch.randn(4, 16)
candidates = torch.randn(4, 5, 16)
output = router(query, candidates)
self.assertEqual(tuple(output["scores"].shape), (4, 5))
self.assertEqual(tuple(output["head_scores"].shape), (4, 5, 2))
self.assertEqual(tuple(output["need_memory_logits"].shape), (4,))
self.assertEqual(tuple(output["hop_logits"].shape), (4, 4))
self.assertEqual(tuple(router.encode_key(query).shape), (4, 8))
self.assertEqual(tuple(router.encode_query(query).shape), (4, 8))
self.assertAlmostEqual(float(router.encode_query(query).norm(dim=-1).mean()), 1.0, places=5)
def test_projected_scores_accepts_a_shared_candidate_bank(self) -> None:
router = _tiny_xl()
query = torch.randn(2, 16)
projected = router.encode_key(torch.randn(7, 16))
scores, head_scores = router.projected_scores(query, projected)
self.assertEqual(tuple(scores.shape), (2, 7))
self.assertEqual(tuple(head_scores.shape), (2, 7, 2))
# A bank-shaped [N, router_dim] key must broadcast like the V2 router.
scores2, _ = router.projected_scores(query, router.encode_key(torch.randn(7, 16)))
self.assertEqual(tuple(scores2.shape), (2, 7))
def test_wrong_shapes_are_rejected(self) -> None:
router = _tiny_xl()
with self.assertRaises(ValueError):
router.encode_key(torch.randn(3, 15))
with self.assertRaises(ValueError):
router.projected_scores(torch.randn(3, 16), torch.randn(3, 5, 7))
def test_arch_config_round_trip_is_exact(self) -> None:
router = _tiny_xl()
config = router.arch_config()
self.assertEqual(config["arch"], ARCH_NAME)
clone = MemoryRouterXL.from_arch_config(config)
self.assertEqual(clone.parameter_count(), router.parameter_count())
self.assertEqual(set(clone.state_dict()), set(router.state_dict()))
clone.load_state_dict(router.state_dict(), strict=True)
x = torch.randn(2, 16)
candidates = torch.randn(2, 6, 16)
router.eval()
clone.eval() # pair dropout is active in train mode; compare deterministically
self.assertTrue(torch.allclose(router(x, candidates)["scores"], clone(x, candidates)["scores"]))
def test_checkpoint_round_trip_via_load_router_xl(self) -> None:
router = _tiny_xl()
with TemporaryDirectory() as directory:
path = Path(directory) / "memory_router_xl.pt"
torch.save({key: value.detach().cpu() for key, value in router.state_dict().items()}, path)
(Path(directory) / "router_arch.json").write_text(
__import__("json").dumps(router.arch_config()), encoding="utf-8"
)
restored = load_router_xl(path)
self.assertEqual(restored.router_dim, router.router_dim)
self.assertEqual(restored.parameter_count(), router.parameter_count())
def test_paged_memory_bank_accepts_the_xl_router(self) -> None:
"""Runtime integration: the fork's bank must drive the new router."""
torch.manual_seed(5)
router = _tiny_xl()
bank = PagedMemoryBankV2(
16,
router=router,
page_capacity=2,
max_pages=64,
hot_pages=2,
top_k_pages=2,
top_k_records=3,
max_hops=3,
coarse_index_bits=8,
)
self.assertEqual(bank.key_dim, router.router_dim)
for index in range(3):
bank.write(
text=f"fact number {index}",
key=torch.randn(16),
entity=f"entity-{index}",
attribute="value",
value=str(index),
confidence=0.9,
)
records, decision = bank.query(query_key=torch.randn(16), query_text="fact number 1")
self.assertTrue(records, "XL router must route at least one active record")
self.assertTrue(decision.record_ids)
self.assertIsInstance(decision.need_memory, bool)
def test_capacity_scales_with_router_dim(self) -> None:
"""Documented capacity jump: XL is a different weight class, not a rename."""
v2 = MemoryRouterV2(2560, router_dim=512, num_heads=8)
v2_params = sum(p.numel() for p in v2.parameters())
xl_small = MemoryRouterXL(2560, router_dim=512, num_heads=16, encoder_hidden=512, pair_hidden=512, policy_hidden=512)
xl_1024 = MemoryRouterXL(2560, router_dim=1024, num_heads=16, encoder_hidden=1024, pair_hidden=1024, policy_hidden=512)
xl_2048 = MemoryRouterXL(2560, router_dim=2048, num_heads=16, encoder_hidden=2048, pair_hidden=2048, policy_hidden=512)
small_params = xl_small.parameter_count()["total"]
mid_params = xl_1024.parameter_count()["total"]
large_params = xl_2048.parameter_count()["total"]
self.assertGreater(small_params, v2_params)
self.assertLess(small_params, mid_params)
self.assertLess(mid_params, large_params)
self.assertEqual(xl_1024.router_dim, 1024)
self.assertEqual(xl_2048.parameter_count()["trainable"], large_params)
if __name__ == "__main__":
unittest.main()