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,68 @@
|
||||
"""Create a deterministic, category-balanced train/eval split for mega memory data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections import Counter, defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--input", required=True)
|
||||
parser.add_argument("--train-output", required=True)
|
||||
parser.add_argument("--eval-output", required=True)
|
||||
parser.add_argument("--eval-per-category", type=int, default=2000)
|
||||
args = parser.parse_args()
|
||||
|
||||
source = Path(args.input)
|
||||
train_path = Path(args.train_output)
|
||||
eval_path = Path(args.eval_output)
|
||||
train_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
eval_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
rows_by_category: dict[str, list[str]] = defaultdict(list)
|
||||
with source.open("r", encoding="utf-8") as handle:
|
||||
for line in handle:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
row = json.loads(line)
|
||||
category = str(row.get("category", "unknown"))
|
||||
rows_by_category[category].append(line)
|
||||
|
||||
if not rows_by_category:
|
||||
raise SystemExit("input contains no rows")
|
||||
if any(len(rows) <= args.eval_per_category for rows in rows_by_category.values()):
|
||||
raise SystemExit("eval-per-category leaves no training rows in at least one category")
|
||||
|
||||
train_counts: Counter[str] = Counter()
|
||||
eval_counts: Counter[str] = Counter()
|
||||
with train_path.open("w", encoding="utf-8") as train_handle, eval_path.open("w", encoding="utf-8") as eval_handle:
|
||||
for category in sorted(rows_by_category):
|
||||
rows = rows_by_category[category]
|
||||
split_at = len(rows) - args.eval_per_category
|
||||
for line in rows[:split_at]:
|
||||
train_handle.write(line + "\n")
|
||||
train_counts[category] += 1
|
||||
for line in rows[split_at:]:
|
||||
eval_handle.write(line + "\n")
|
||||
eval_counts[category] += 1
|
||||
|
||||
manifest = {
|
||||
"source": str(source),
|
||||
"eval_per_category": args.eval_per_category,
|
||||
"categories": sorted(rows_by_category),
|
||||
"train_counts": dict(train_counts),
|
||||
"eval_counts": dict(eval_counts),
|
||||
"train_rows": sum(train_counts.values()),
|
||||
"eval_rows": sum(eval_counts.values()),
|
||||
}
|
||||
manifest_path = train_path.parent / "mega_split_manifest.json"
|
||||
manifest_path.write_text(json.dumps(manifest, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(json.dumps(manifest, ensure_ascii=False, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user