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,109 @@
|
||||
"""Package learned memory into an adapter checkpoint and verify restart/reset."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from .evaluate_native_memory import _controller_step, _encode_prompt
|
||||
from .qwen_integration import (
|
||||
DEFAULT_MEMORY_RESET_TOKEN,
|
||||
load_memory_config,
|
||||
load_qwen_dynamic,
|
||||
load_tokenizer,
|
||||
resolve_memory_reset_token,
|
||||
)
|
||||
from .train_native_memory import load_records
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default=".")
|
||||
parser.add_argument("--adapter", default="V2_dpskw/qwen_memory_adapter_native_v3")
|
||||
parser.add_argument("--data", default="V2_dpskw/data/native_memory/eval.jsonl")
|
||||
parser.add_argument(
|
||||
"--output-adapter",
|
||||
default="V2_dpskw/qwen_memory_adapter_native_v3_persistent",
|
||||
)
|
||||
parser.add_argument("--max-length", type=int, default=192)
|
||||
parser.add_argument("--max-new-tokens", type=int, default=8)
|
||||
parser.add_argument("--no-4bit", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
records = load_records(args.data)
|
||||
record = next(record for record in records if record.get("answerable"))
|
||||
tokenizer = load_tokenizer(args.model_path)
|
||||
config = load_memory_config(args.adapter)
|
||||
config.persistent_memory = True
|
||||
config.reset_token_id = resolve_memory_reset_token(tokenizer, DEFAULT_MEMORY_RESET_TOKEN)
|
||||
model = load_qwen_dynamic(args.model_path, memory_config=config, load_in_4bit=not args.no_4bit)
|
||||
model.load_memory_adapter(args.adapter)
|
||||
model.reset_memory()
|
||||
for chunk in record["memory_chunks"]:
|
||||
_controller_step(model, tokenizer, chunk, args.max_length)
|
||||
output_adapter = Path(args.output_adapter)
|
||||
model.save_persistent_memory_checkpoint(output_adapter)
|
||||
saved_norm = float(model.runtime.state.detach().float().norm())
|
||||
del model
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
restart_config = load_memory_config(output_adapter)
|
||||
restarted = load_qwen_dynamic(
|
||||
args.model_path,
|
||||
memory_config=restart_config,
|
||||
load_in_4bit=not args.no_4bit,
|
||||
)
|
||||
restarted.load_memory_adapter(output_adapter)
|
||||
answer = str(record["answer"])
|
||||
prompt = _encode_prompt(tokenizer, record["query"][:-1])
|
||||
device = restarted._find_layer_device()
|
||||
generated = restarted.generate(
|
||||
input_ids=prompt.to(device),
|
||||
attention_mask=torch.ones_like(prompt, device=device),
|
||||
max_new_tokens=args.max_new_tokens,
|
||||
do_sample=False,
|
||||
use_cache=False,
|
||||
update_memory=False,
|
||||
)
|
||||
response = tokenizer.decode(
|
||||
generated[:, prompt.shape[1] :][0].detach().cpu().tolist(),
|
||||
skip_special_tokens=True,
|
||||
).strip()
|
||||
reset_prompt = _encode_prompt(tokenizer, [{"role": "user", "content": DEFAULT_MEMORY_RESET_TOKEN}])
|
||||
restarted.generate(
|
||||
input_ids=reset_prompt.to(device),
|
||||
attention_mask=torch.ones_like(reset_prompt, device=device),
|
||||
max_new_tokens=1,
|
||||
do_sample=False,
|
||||
use_cache=False,
|
||||
update_memory=False,
|
||||
)
|
||||
reset_norm = float(restarted.runtime.state.detach().float().norm())
|
||||
report = {
|
||||
"source_adapter": str(args.adapter),
|
||||
"output_adapter": str(output_adapter),
|
||||
"record_id": record.get("id"),
|
||||
"answer": answer,
|
||||
"generated_after_restart": response,
|
||||
"restart_contains_answer": answer in response,
|
||||
"saved_memory_norm": saved_norm,
|
||||
"reset_token": DEFAULT_MEMORY_RESET_TOKEN,
|
||||
"reset_token_id": restart_config.reset_token_id,
|
||||
"memory_norm_after_reset_token": reset_norm,
|
||||
"reset_zeroed_memory": reset_norm < 1e-5,
|
||||
}
|
||||
(output_adapter / "native_checkpoint_verification.json").write_text(
|
||||
json.dumps(report, ensure_ascii=False, indent=2),
|
||||
encoding="utf-8",
|
||||
)
|
||||
print(json.dumps(report, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user