"""Paired teacher/student benchmark for Natural Memory versus full KV context.""" from __future__ import annotations import argparse import gc import json import re import sys import time from collections import defaultdict from pathlib import Path from typing import Any, Iterable import torch from .qwen_integration import load_memory_config, load_qwen_base, load_qwen_dynamic, load_tokenizer from .stream_chat_qwen_memory import _chat_tensor, _write_turn PROJECT_ROOT = Path(__file__).resolve().parent def _path(value: str | Path) -> Path: path = Path(value) if path.is_absolute() or path.exists(): return path return PROJECT_ROOT / path def _normalize(text: str) -> str: return re.sub(r"[\s`*_#,。!?、;:,.!?;:'\"()()\[\]{}]", "", str(text)).lower() def _contains(text: str, choices: Iterable[str]) -> bool: normalized = _normalize(text) return any(_normalize(choice) and _normalize(choice) in normalized for choice in choices) def _passed(response: str, case: dict[str, Any]) -> bool: return _contains(response, case.get("acceptable", [])) and not _contains( response, case.get("forbidden", []) ) def _read_cases(path: Path, *, limit: int | None, offset: int, category: str | None) -> list[dict[str, Any]]: rows: list[dict[str, Any]] = [] with path.open("r", encoding="utf-8") as handle: for raw in handle: if not raw.strip(): continue case = json.loads(raw) if category and case.get("category") != category: continue if offset > 0: offset -= 1 continue rows.append(case) if limit is not None and len(rows) >= limit: break if not rows: raise ValueError("no validation cases selected") return rows def _teacher_messages(case: dict[str, Any]) -> list[dict[str, str]]: messages: list[dict[str, str]] = [] for fact in case["facts"]: messages.append({"role": "user", "content": str(fact["text"])}) messages.append({"role": "assistant", "content": str(fact.get("assistant", "好的。"))}) messages.append({"role": "user", "content": str(case["query"])}) return messages @torch.inference_mode() def _generate_teacher(model: Any, tokenizer: Any, case: dict[str, Any], max_new_tokens: int) -> str: device = model.get_input_embeddings().weight.device encoded = tokenizer.apply_chat_template( _teacher_messages(case), tokenize=True, add_generation_prompt=True, return_tensors="pt", return_dict=True, enable_thinking=False, ) encoded = {key: value.to(device) for key, value in encoded.items() if isinstance(value, torch.Tensor)} output = model.generate( **encoded, max_new_tokens=max_new_tokens, do_sample=False, use_cache=True, pad_token_id=tokenizer.pad_token_id, ) return tokenizer.decode(output[0, encoded["input_ids"].shape[1] :], skip_special_tokens=True).strip() @torch.inference_mode() def _generate_student( model: Any, tokenizer: Any, case: dict[str, Any], *, max_new_tokens: int, force_write: bool, ) -> tuple[str, dict[str, Any]]: device = model._find_layer_device() model.reset_memory(batch_size=1, device=device) write_rows: list[dict[str, Any]] = [] for fact in case["facts"]: changed = _write_turn( model, tokenizer, str(fact["text"]), device, force_write=force_write, ) last_slot = model.runtime.text_last_written_slot write_rows.append( { "kind": fact.get("kind", "fact"), "should_write": bool(fact.get("should_write", True)), "changed": bool(changed), "slot": int(last_slot[0].item()) if isinstance(last_slot, torch.Tensor) else -1, } ) encoded = {key: value.to(device) for key, value in _chat_tensor( tokenizer, str(case["query"]), ).items()} query = tokenizer(str(case["query"]), add_special_tokens=False, return_tensors="pt") query_ids = query["input_ids"].to(device) query_mask = query.get("attention_mask") if query_mask is None: query_mask = torch.ones_like(query_ids) query_mask = query_mask.to(device) output = model.generate( **encoded, max_new_tokens=max_new_tokens, do_sample=False, update_memory=False, memory_query_input_ids=query_ids, memory_query_attention_mask=query_mask, memory_query_text=str(case["query"]), use_cache=True, pad_token_id=tokenizer.pad_token_id, ) response = tokenizer.decode(output[0, encoded["input_ids"].shape[1] :], skip_special_tokens=True).strip() return response, { "writes": write_rows, "valid_slots": int(model.runtime.text_slot_valid.sum().item()) if isinstance(model.runtime.text_slot_valid, torch.Tensor) else 0, "v2": model.memory_v2_stats(), } def _release(model: Any) -> None: del model gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() def _summarize(rows: list[dict[str, Any]]) -> dict[str, Any]: categories: dict[str, list[float]] = defaultdict(list) for row in rows: categories[str(row["category"])].append(float(row["passed"])) return { "cases": len(rows), "accuracy": sum(float(row["passed"]) for row in rows) / max(1, len(rows)), "categories": { category: {"cases": len(values), "accuracy": sum(values) / len(values)} for category, values in sorted(categories.items()) }, } def main() -> None: if hasattr(sys.stdout, "reconfigure"): sys.stdout.reconfigure(encoding="utf-8", errors="replace") parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--model-path", default="qwen3_5_4b_natural_memory_v2") parser.add_argument("--adapter", default=None, help="optional candidate adapter; omitted uses embedded policy") parser.add_argument("--data", default="data/mega_validation/memory_validation_100k.jsonl") parser.add_argument("--output", default="mega_memory_vs_full_kv_report.json") parser.add_argument("--limit", type=int, default=32) parser.add_argument("--offset", type=int, default=0) parser.add_argument("--category", default=None) parser.add_argument("--max-new-tokens", type=int, default=24) parser.add_argument("--force-write", action="store_true") parser.add_argument("--no-4bit", action="store_true") args = parser.parse_args() model_path = _path(args.model_path) data_path = _path(args.data) output_path = _path(args.output) cases = _read_cases(data_path, limit=args.limit, offset=args.offset, category=args.category) tokenizer = load_tokenizer(model_path) use_4bit = not args.no_4bit report: dict[str, Any] = { "format_version": 1, "model_path": str(model_path), "data": str(data_path), "selected_cases": len(cases), "force_write": bool(args.force_write), "quantization": "4bit_nf4" if use_4bit else "none", } started = time.perf_counter() teacher = load_qwen_base(model_path, load_in_4bit=use_4bit) teacher.eval() teacher_rows: list[dict[str, Any]] = [] for index, case in enumerate(cases, 1): response = _generate_teacher(teacher, tokenizer, case, args.max_new_tokens) teacher_rows.append( { "id": case["id"], "category": case["category"], "response": response, "passed": _passed(response, case), } ) if index % 16 == 0: print(f"teacher {index}/{len(cases)}") report["teacher"] = _summarize(teacher_rows) _release(teacher) adapter_path = _path(args.adapter) if args.adapter else None config_source = adapter_path or model_path config = load_memory_config(config_source) config.persistent_memory = True config.natural_language_memory = True config.automatic_memory = True student = load_qwen_dynamic(model_path, memory_config=config, load_in_4bit=use_4bit) if adapter_path is not None: student.load_memory_adapter(adapter_path, strict=True) student.eval() student_rows: list[dict[str, Any]] = [] for index, case in enumerate(cases, 1): response, diagnostics = _generate_student( student, tokenizer, case, max_new_tokens=args.max_new_tokens, force_write=args.force_write, ) student_rows.append( { "id": case["id"], "category": case["category"], "response": response, "passed": _passed(response, case), "diagnostics": diagnostics, } ) if index % 16 == 0: print(f"student {index}/{len(cases)}") report["student"] = _summarize(student_rows) _release(student) teacher_accuracy = float(report["teacher"]["accuracy"]) student_accuracy = float(report["student"]["accuracy"]) teacher_by_id = {row["id"]: row for row in teacher_rows} paired_teacher_pass = sum(bool(teacher_by_id[row["id"]]["passed"]) for row in student_rows) paired_student_pass = sum(bool(row["passed"]) and teacher_by_id[row["id"]]["passed"] for row in student_rows) category_gate: dict[str, Any] = {} for category in sorted({str(case["category"]) for case in cases}): teacher_cat = [row for row in teacher_rows if row["category"] == category] student_cat = [row for row in student_rows if row["category"] == category] t = sum(float(row["passed"]) for row in teacher_cat) / max(1, len(teacher_cat)) s = sum(float(row["passed"]) for row in student_cat) / max(1, len(student_cat)) category_gate[category] = {"teacher_accuracy": t, "student_accuracy": s, "ratio": s / max(t, 1e-9), "pass": s >= 0.95 * t} report["parity"] = { "teacher_accuracy": teacher_accuracy, "student_accuracy": student_accuracy, "student_to_teacher_ratio": student_accuracy / max(teacher_accuracy, 1e-9), "paired_teacher_pass": paired_teacher_pass, "paired_student_pass": paired_student_pass, "paired_ratio": paired_student_pass / max(1, paired_teacher_pass), "category_gate": category_gate, "required_ratio": 0.95, "pass": student_accuracy >= 0.95 * teacher_accuracy and all(item["pass"] for item in category_gate.values()), } report["elapsed_seconds"] = time.perf_counter() - started report["failures"] = [row for row in student_rows if not row["passed"]][:100] output_path.parent.mkdir(parents=True, exist_ok=True) output_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding="utf-8") print(json.dumps({"teacher": report["teacher"], "student": report["student"], "parity": report["parity"], "output": str(output_path)}, ensure_ascii=False, indent=2)) if __name__ == "__main__": main()