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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user