item 32 pre-registered (Nils's synthesis): discrete latent chain — lens-snapped token embeddings fed back via zero-init projector (ST top-32, TF/free-running arms, frozen arm-1 merge); sym_iterate in carry_common, trainer/eval wiring

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Nils
2026-07-18 00:16:08 +02:00
co-authored by Claude Fable 5
parent 3db6dd8fef
commit f0b7942c7a
5 changed files with 193 additions and 4 deletions
+27 -1
View File
@@ -43,6 +43,9 @@ def main():
ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT",
help="KVMemoryAdapter checkpoint: burst states become "
"band-layer KV prefix entries at generation")
ap.add_argument("--symchain", default=None, metavar="PROJ_PT",
help="item 32: discrete latent chain — zero-init "
"projector checkpoint; generation snaps hard argmax")
args = ap.parse_args()
model, tok = load_model(dtype=torch.bfloat16)
@@ -71,6 +74,29 @@ def main():
print(f"band-lora loaded: {args.bandlora} "
f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})",
flush=True)
symchain = None
if args.symchain:
d_ = model.config.get_text_config().hidden_size
proj = torch.nn.Linear(d_, d_).cuda()
proj.load_state_dict(torch.load(args.symchain, map_location="cuda"))
proj.eval()
jbar = torch.load(Path(__file__).resolve().parent.parent
/ "results/jbar.pt", map_location="cuda")["Jbar"]
from loop_common import BAND
J30 = jbar[BAND[1]].float()
tm = model.model.language_model
softcap = model.config.get_text_config().final_logit_softcapping
def lens_fn(h):
x = tm.norm((h.float() @ J30.T).to(tm.norm.weight.dtype))
lg = model.lm_head(x)
return softcap * torch.tanh(lg / softcap) if softcap else lg
symchain = {"proj": proj, "lens_fn": lens_fn,
"embed_w": model.get_input_embeddings().weight.detach(),
"start_id": tok("\n",
add_special_tokens=False)["input_ids"][0]}
print(f"symchain loaded: {args.symchain}", flush=True)
adapter = MergeAdapter(
d=model.config.get_text_config().hidden_size).cuda()
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
@@ -101,7 +127,7 @@ def main():
max_new_tokens=args.max_new,
feedforward=args.feedforward,
inner_iters=args.inner_iters,
kvmem=kvmem)
kvmem=kvmem, symchain=symchain)
if kvmem is not None:
from kv_memory import arm_memory
arm_memory(None)