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