Files
natural-memory-nm21/train_auto_memory_policy.py
T
WpyQwq 643e22ecb9 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,读写关闭时与原生模型逐位相同
2026-09-19 11:11:31 +08:00

285 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Train the internal high-recall automatic memory write policy.
The frozen Qwen backbone and the existing native memory controller provide
the representation. Only a small policy head is trained. Positive examples
cover durable personal facts, preferences, plans, project constraints and
corrections; negative examples cover questions, requests, hypotheticals and
casual conversation. The runtime still stores the exact user token sequence,
so this head decides *whether* to remember rather than compressing the fact.
"""
from __future__ import annotations
import argparse
import json
import random
from pathlib import Path
import torch
import torch.nn.functional as F
from .qwen_integration import load_memory_config, load_qwen_dynamic, load_tokenizer
POSITIVE = (
"我叫{value}。",
"我的常住城市是{value}。",
"我最喜欢的水果是{value}。",
"我正在开发{value}项目。",
"以后请把代码默认写成{value}。",
"我计划在{value}完成这个任务。",
"我的工作地点是{value}。",
"我的常用时区是{value}。",
"请记住这条信息:{value}。",
"更正一下,刚才的内容应该是{value}。",
"这是我的长期偏好:{value}。",
"这个项目的重要约束是{value}。",
)
NEGATIVE = (
"帮我写一段关于{value}的代码。",
"请解释{value}是什么意思。",
"{value}是什么?",
"你觉得{value}怎么样?",
"如果以后遇到{value},应该怎么办?",
"今天天气不错,随便聊聊{value}。",
"请把{value}翻译成英文。",
"计算一下{value}。",
"给我介绍一下{value}。",
"哈哈,{value}真有意思。",
"我想知道之前有没有提到{value}。",
"假设我选择{value},会发生什么?",
"我叫什么?",
"我的名字是什么?",
"你还记得我叫什么吗?",
"我正在开发什么项目?",
"我的项目叫什么?",
"请问我的项目叫什么?",
"我的工作地点是什么?",
"我的工作地点代号是什么?",
"工作地点代号是多少?",
"你记得我的工作地点吗?",
"请告诉我之前有没有说过{value}。",
"我之前有没有告诉过你{value}?",
"能不能帮我完成{value}?",
"如何处理{value}?",
"请给我一个{value}的方案。",
)
VALUES = (
"小明",
"上海",
"红富士苹果",
"个人记忆系统",
"Python",
"下周五",
"R7",
"Asia/Shanghai",
"不要删除用户数据",
"使用简洁中文",
"每天晚上八点",
"蓝鲸-47",
)
def make_examples(seed: int, count: int) -> list[tuple[str, float]]:
rng = random.Random(seed)
examples: list[tuple[str, float]] = []
for _ in range(count):
examples.append((rng.choice(POSITIVE).format(value=rng.choice(VALUES)), 1.0))
examples.append((rng.choice(NEGATIVE).format(value=rng.choice(VALUES)), 0.0))
rng.shuffle(examples)
return examples
def encode_batch(tokenizer, texts: list[str], device: torch.device):
rows = []
masks = []
for text in texts:
encoded = tokenizer.apply_chat_template(
[{"role": "user", "content": text}],
tokenize=True,
add_generation_prompt=True,
return_tensors="pt",
return_dict=True,
enable_thinking=False,
)
rows.append(encoded["input_ids"][0])
masks.append(encoded.get("attention_mask", torch.ones_like(encoded["input_ids"]))[0])
max_length = max(row.numel() for row in rows)
input_ids = torch.zeros(len(rows), max_length, dtype=torch.long, device=device)
attention_mask = torch.zeros(len(rows), max_length, dtype=torch.long, device=device)
for index, (row, mask) in enumerate(zip(rows, masks)):
input_ids[index, : row.numel()] = row.to(device)
attention_mask[index, : mask.numel()] = mask.to(device)
return input_ids, attention_mask
@torch.inference_mode()
def collect_representations(model, tokenizer, examples, *, batch_size: int, device):
vectors = []
labels = []
for start in range(0, len(examples), batch_size):
batch = examples[start : start + batch_size]
input_ids, attention_mask = encode_batch(
tokenizer,
[item[0] for item in batch],
device,
)
model.reset_memory(batch_size=len(batch), device=device)
model(
input_ids=input_ids,
attention_mask=attention_mask,
read_memory=False,
update_memory=True,
return_memory=True,
use_cache=False,
)
representation = getattr(model.memory, "last_write_representation", None)
if representation is None:
raise RuntimeError("native controller did not expose a write representation")
vectors.append(representation.detach().float().cpu())
labels.extend(item[1] for item in batch)
return torch.cat(vectors), torch.tensor(labels, dtype=torch.float32)
def evaluate(policy, vectors, labels, threshold: float) -> dict[str, float]:
with torch.inference_mode():
probabilities = torch.sigmoid(policy(vectors)).cpu()
labels = labels.cpu()
predictions = probabilities >= threshold
positive = labels >= 0.5
negative = ~positive
true_positive = (predictions & positive).sum().item()
false_negative = ((~predictions) & positive).sum().item()
false_positive = (predictions & negative).sum().item()
true_negative = ((~predictions) & negative).sum().item()
return {
"threshold": threshold,
"accuracy": float((predictions == positive).float().mean()),
"positive_recall": true_positive / max(1, true_positive + false_negative),
"negative_specificity": true_negative / max(1, true_negative + false_positive),
"false_positive_rate": false_positive / max(1, false_positive + true_negative),
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default=".")
parser.add_argument(
"--base-adapter",
default="V2_dpskw/qwen_memory_adapter_natural_controller_v3",
)
parser.add_argument(
"--output-adapter",
default="V2_dpskw/qwen_memory_adapter_natural_auto_v3",
)
parser.add_argument("--steps", type=int, default=2400)
parser.add_argument("--example-count", type=int, default=1280)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--lr", type=float, default=2e-4)
parser.add_argument("--threshold", type=float, default=0.35)
parser.add_argument(
"--text-memory-threshold",
type=float,
default=0.30,
help="retrieval threshold used by the automatic-memory adapter",
)
parser.add_argument("--seed", type=int, default=20260904)
parser.add_argument("--no-4bit", action="store_true")
args = parser.parse_args()
random.seed(args.seed)
torch.manual_seed(args.seed)
tokenizer = load_tokenizer(args.model_path)
config = load_memory_config(args.base_adapter)
config.natural_language_memory = True
config.automatic_memory = True
config.auto_memory_threshold = args.threshold
config.text_memory_threshold = args.text_memory_threshold
config.persistent_memory = False
model = load_qwen_dynamic(
args.model_path,
memory_config=config,
load_in_4bit=not args.no_4bit,
)
model.load_memory_adapter(args.base_adapter, strict=True)
if model.memory_policy is None:
raise RuntimeError("automatic memory policy was not created")
model.eval()
model.memory_policy.train()
device = model._find_layer_device()
examples = make_examples(args.seed, args.example_count)
split = int(len(examples) * 0.8)
train_examples = examples[:split]
eval_examples = examples[split:]
train_vectors, train_labels = collect_representations(
model,
tokenizer,
train_examples,
batch_size=args.batch_size,
device=device,
)
eval_vectors, eval_labels = collect_representations(
model,
tokenizer,
eval_examples,
batch_size=args.batch_size,
device=device,
)
train_vectors = train_vectors.to(device)
train_labels = train_labels.to(device)
eval_vectors = eval_vectors.to(device)
eval_labels = eval_labels.to(device)
optimizer = torch.optim.AdamW(model.memory_policy.parameters(), lr=args.lr, weight_decay=0.01)
positive_weight = torch.tensor([1.5], device=device)
rng = random.Random(args.seed + 1)
for step in range(1, args.steps + 1):
indices = torch.tensor(
[rng.randrange(train_vectors.shape[0]) for _ in range(args.batch_size)],
dtype=torch.long,
device=device,
)
logits = model.memory_policy(train_vectors[indices])
weights = torch.where(train_labels[indices] >= 0.5, positive_weight, torch.ones_like(logits))
loss = F.binary_cross_entropy_with_logits(logits, train_labels[indices], weight=weights)
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.memory_policy.parameters(), 1.0)
optimizer.step()
if step == 1 or step % 100 == 0 or step == args.steps:
stats = evaluate(model.memory_policy, train_vectors, train_labels, args.threshold)
print(
f"step={step} loss={float(loss.detach()):.5f} "
f"train_accuracy={stats['accuracy']:.3f} "
f"positive_recall={stats['positive_recall']:.3f} "
f"false_positive_rate={stats['false_positive_rate']:.3f}"
)
model.memory_policy.eval()
model._memory_policy_ready = True
output_dir = Path(args.output_adapter)
model.save_memory_adapter(output_dir)
report = {
"steps": args.steps,
"example_count_per_class": args.example_count,
"train_examples": len(train_examples),
"eval_examples": len(eval_examples),
"source_adapter": str(args.base_adapter),
"policy": "high_recall_automatic_memory_importance",
"threshold": args.threshold,
"train": evaluate(model.memory_policy, train_vectors, train_labels, args.threshold),
"eval": evaluate(model.memory_policy, eval_vectors, eval_labels, args.threshold),
}
(output_dir / "auto_policy_training.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()