- 引入 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,读写关闭时与原生模型逐位相同
108 lines
3.4 KiB
Python
108 lines
3.4 KiB
Python
"""Verify answering after a model restart using only a saved user memory state."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import gc
|
|
import json
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
|
|
from .benchmark_qwen import _generation_prompt
|
|
from .qwen_integration import load_memory_config, load_qwen_dynamic, load_tokenizer
|
|
from .train_qwen_memory import encode_messages, pad_batch
|
|
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--model-path", default=".")
|
|
parser.add_argument(
|
|
"--adapter",
|
|
default="V2_dpskw/qwen_memory_adapter_pointer",
|
|
)
|
|
parser.add_argument(
|
|
"--data",
|
|
default="V2_dpskw/data/benchmark_eval.jsonl",
|
|
)
|
|
parser.add_argument("--record-index", type=int, default=0)
|
|
parser.add_argument(
|
|
"--memory-state",
|
|
default="V2_dpskw/data/user_memory_demo.pt",
|
|
)
|
|
parser.add_argument("--max-length", type=int, default=128)
|
|
parser.add_argument("--max-new-tokens", type=int, default=4)
|
|
parser.add_argument("--no-4bit", action="store_true")
|
|
args = parser.parse_args()
|
|
|
|
records = [
|
|
json.loads(line)
|
|
for line in Path(args.data).read_text(encoding="utf-8").splitlines()
|
|
if line.strip()
|
|
]
|
|
record = records[args.record_index]
|
|
tokenizer = load_tokenizer(args.model_path)
|
|
memory_config = load_memory_config(args.adapter)
|
|
|
|
model = load_qwen_dynamic(
|
|
args.model_path,
|
|
memory_config=memory_config,
|
|
load_in_4bit=not args.no_4bit,
|
|
)
|
|
model.load_memory_adapter(args.adapter)
|
|
model.eval()
|
|
device = model._find_layer_device()
|
|
memory = encode_messages(tokenizer, record["memory"], args.max_length)
|
|
memory_input, memory_mask, _ = pad_batch([memory], int(tokenizer.pad_token_id))
|
|
state_path = Path(args.memory_state)
|
|
|
|
with torch.inference_mode():
|
|
model(
|
|
input_ids=memory_input.to(device),
|
|
attention_mask=memory_mask.to(device),
|
|
read_memory=False,
|
|
update_memory=True,
|
|
return_memory=True,
|
|
use_cache=False,
|
|
)
|
|
model.save_runtime_memory(state_path)
|
|
print(f"saved_memory_state={state_path}")
|
|
|
|
del model
|
|
gc.collect()
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|
|
|
|
restarted = load_qwen_dynamic(
|
|
args.model_path,
|
|
memory_config=memory_config,
|
|
load_in_4bit=not args.no_4bit,
|
|
)
|
|
restarted.load_memory_adapter(args.adapter)
|
|
restarted.eval()
|
|
device = restarted._find_layer_device()
|
|
restarted.load_runtime_memory(state_path, device=device)
|
|
prompt = _generation_prompt(tokenizer, record["query"][:-1], device)
|
|
with torch.inference_mode():
|
|
output = restarted.generate(
|
|
**prompt,
|
|
max_new_tokens=args.max_new_tokens,
|
|
do_sample=False,
|
|
update_memory=False,
|
|
use_cache=True,
|
|
pad_token_id=tokenizer.pad_token_id,
|
|
)
|
|
generated = tokenizer.decode(
|
|
output[0, prompt["input_ids"].shape[1] :],
|
|
skip_special_tokens=True,
|
|
).replace(" ", "").replace("\r", "").replace("\n", "").strip()
|
|
expected = str(record["answer"])
|
|
print(f"query_contains_history=false")
|
|
print(f"expected={expected}")
|
|
print(f"restarted_generated={generated}")
|
|
print(f"correct={generated.startswith(expected)}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|