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:
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user