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:
@@ -1038,3 +1038,38 @@ to explore any new channel; a headroom-bearing task would be a fairer
|
|||||||
test. Verdict as registered: within +-5 -> no evidence that a native
|
test. Verdict as registered: within +-5 -> no evidence that a native
|
||||||
attention read path breaches the consumption wall, with the init
|
attention read path breaches the consumption wall, with the init
|
||||||
confound flagged as the one loose thread.
|
confound flagged as the one loose thread.
|
||||||
|
|
||||||
|
32. **The discrete latent chain — "latent paper" (pre-registered
|
||||||
|
2026-07-18 ~00:50, before running; Nils's synthesis: "the loop
|
||||||
|
needs paper, but we don't want that to be full tokens but still
|
||||||
|
latent space").** Diagnosis from the full matrix: every latent
|
||||||
|
medium lacked DISCRETENESS — tokens' magic is the snap
|
||||||
|
(error-correction per step), not visibility. The lens is a native
|
||||||
|
codebook: argmax over its readout quantizes any band state onto
|
||||||
|
the model's own symbol space. Design (each burst tick):
|
||||||
|
read s_{i-1} through the frozen lens; snap to a token
|
||||||
|
(straight-through over top-32, hard forward / soft gradient; tick
|
||||||
|
0 = newline "a step begins"); feed E(token) back through a
|
||||||
|
ZERO-INIT projector alongside the analog carry:
|
||||||
|
x_i = merge(e, s_{i-1}) + proj(E(sym)). Two rails, matching the
|
||||||
|
microscopy's own two-channel algorithm: analog (plans/magnitudes)
|
||||||
|
+ discrete (exact symbols). Writing and computing coincide by
|
||||||
|
construction: what the lens reads IS what gets transported —
|
||||||
|
item 25's lce loss (reused, λ=0.3) is now load-bearing, and its
|
||||||
|
proven concentration effect (10.3->1.9) doubles as the
|
||||||
|
quantization pressure that collapses the diffuse thinking-state
|
||||||
|
superposition (measured: pause states spread over ~hundreds of
|
||||||
|
tokens, committed states 1-2). Merge FROZEN at item-29 arm-1
|
||||||
|
(39.1); only the 2.4M projector trains — the by-construction
|
||||||
|
ablation is 39.1. Arms: (a) sctf — teacher-forced symbols
|
||||||
|
(ground-truth deleted-step tokens; item-29's winning recipe);
|
||||||
|
(b) scst — free-running straight-through snaps. Eval always
|
||||||
|
hard-argmax free-running: 0:0 sanity + 2:0, n=256. Decision vs
|
||||||
|
39.1: >44.1 = discreteness was the missing paper property, the
|
||||||
|
latent-chain program reopens (widenings pre-sketched: top-k
|
||||||
|
parallel snaps, L22/L26 depth rails, tape via slots/kvmem);
|
||||||
|
within +-5 = the discrete channel adds nothing over the analog
|
||||||
|
carry and the "loop = plan machine, tokens = executor" division
|
||||||
|
stands as final. Smokes: TF/ST train (loss arithmetic exact,
|
||||||
|
warm-start vals intact), eval generates with hard snaps.
|
||||||
|
Job: scripts/jobs/zzz_y_symchain.sh.
|
||||||
|
|||||||
+47
-2
@@ -89,6 +89,42 @@ def splice_inner_iters(updates, inner_iters, inner_at, prompt_lens, dev, B):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def sym_iterate(looper, adapter, proj, e, calls, S, X, rows, anchor, m,
|
||||||
|
lens_fn, embed_w, sym_tf=None, start_id=None, topk=32,
|
||||||
|
use_checkpoint=False, iter_states=None):
|
||||||
|
"""Item 32: discrete latent chain at the anchor. Each tick reads the
|
||||||
|
previous anchor state through the lens, snaps it to a token
|
||||||
|
(straight-through over top-k) or takes the teacher token (sym_tf:
|
||||||
|
(B, m) ids, teacher forcing), and feeds that token's embedding back
|
||||||
|
through a zero-init projector ALONGSIDE the analog carry:
|
||||||
|
x_i = merge(e, s_{i-1}) + proj(E(sym))
|
||||||
|
Tick 0 uses start_id (a newline: 'a step begins')."""
|
||||||
|
for i in range(m):
|
||||||
|
s_prev = S[rows, anchor]
|
||||||
|
if sym_tf is not None:
|
||||||
|
symb = embed_w[sym_tf[:, i]]
|
||||||
|
elif i == 0:
|
||||||
|
symb = embed_w[torch.full((rows.shape[0],), start_id,
|
||||||
|
device=e.device)]
|
||||||
|
else:
|
||||||
|
logits = lens_fn(s_prev).float()
|
||||||
|
p, idx = torch.softmax(logits, -1).topk(topk, dim=-1)
|
||||||
|
p = p / p.sum(-1, keepdim=True)
|
||||||
|
soft = (p.unsqueeze(-1) * embed_w[idx].float()).sum(-2)
|
||||||
|
hard = embed_w[idx[:, 0]].float()
|
||||||
|
symb = hard + soft - soft.detach()
|
||||||
|
x_new = (adapter(e[rows, anchor], s_prev).float()
|
||||||
|
+ proj(symb.float()))
|
||||||
|
X = X.clone()
|
||||||
|
X[rows, anchor] = x_new.to(X.dtype)
|
||||||
|
S = (checkpoint(lambda X_: looper.band(X_, calls), X,
|
||||||
|
use_reentrant=False) if use_checkpoint
|
||||||
|
else looper.band(X, calls))
|
||||||
|
if iter_states is not None:
|
||||||
|
iter_states.append(S[rows, anchor])
|
||||||
|
return S, X
|
||||||
|
|
||||||
|
|
||||||
def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
|
def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
|
||||||
k, use_checkpoint=False, feedforward=False,
|
k, use_checkpoint=False, feedforward=False,
|
||||||
return_states=False, inner_iters=0, inner_at=None,
|
return_states=False, inner_iters=0, inner_at=None,
|
||||||
@@ -128,7 +164,7 @@ def carry_logits(looper, adapter, input_ids, attention_mask, prompt_lens,
|
|||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def generate_carry_c(looper, adapter, tok, input_ids, attention_mask,
|
def generate_carry_c(looper, adapter, tok, input_ids, attention_mask,
|
||||||
k, p, max_new_tokens=10, feedforward=False,
|
k, p, max_new_tokens=10, feedforward=False,
|
||||||
inner_iters=0, kvmem=None):
|
inner_iters=0, kvmem=None, symchain=None):
|
||||||
"""Greedy design-C generation (left-padded batch, uniform positions).
|
"""Greedy design-C generation (left-padded batch, uniform positions).
|
||||||
|
|
||||||
Appends p pause tokens, prefill-loops the prompt, carries through the
|
Appends p pause tokens, prefill-loops the prompt, carries through the
|
||||||
@@ -159,7 +195,16 @@ def generate_carry_c(looper, adapter, tok, input_ids, attention_mask,
|
|||||||
updates = [(torch.arange(B, device=dev),
|
updates = [(torch.arange(B, device=dev),
|
||||||
torch.full((B,), n_prompt + j, device=dev,
|
torch.full((B,), n_prompt + j, device=dev,
|
||||||
dtype=torch.long)) for j in range(p)]
|
dtype=torch.long)) for j in range(p)]
|
||||||
if inner_iters:
|
if symchain is not None and inner_iters:
|
||||||
|
anchor_sc = torch.full((B,), n_prompt + p - 1, device=dev,
|
||||||
|
dtype=torch.long)
|
||||||
|
S, X = sym_iterate(
|
||||||
|
looper, adapter, symchain["proj"], e, calls, S, X,
|
||||||
|
torch.arange(B, device=dev), anchor_sc, inner_iters,
|
||||||
|
symchain["lens_fn"], symchain["embed_w"],
|
||||||
|
start_id=symchain["start_id"])
|
||||||
|
updates = updates # pauses (if any) already handled above
|
||||||
|
elif inner_iters:
|
||||||
anchor = torch.full((B,), n_prompt + p - 1, device=dev,
|
anchor = torch.full((B,), n_prompt + p - 1, device=dev,
|
||||||
dtype=torch.long)
|
dtype=torch.long)
|
||||||
inner = [(torch.arange(B, device=dev), anchor, True)
|
inner = [(torch.arange(B, device=dev), anchor, True)
|
||||||
|
|||||||
@@ -43,6 +43,9 @@ def main():
|
|||||||
ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT",
|
ap.add_argument("--kvmem", default=None, metavar="KVMEM_PT",
|
||||||
help="KVMemoryAdapter checkpoint: burst states become "
|
help="KVMemoryAdapter checkpoint: burst states become "
|
||||||
"band-layer KV prefix entries at generation")
|
"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()
|
args = ap.parse_args()
|
||||||
|
|
||||||
model, tok = load_model(dtype=torch.bfloat16)
|
model, tok = load_model(dtype=torch.bfloat16)
|
||||||
@@ -71,6 +74,29 @@ def main():
|
|||||||
print(f"band-lora loaded: {args.bandlora} "
|
print(f"band-lora loaded: {args.bandlora} "
|
||||||
f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})",
|
f"(r={ck['rank']}, layers {ck['band'][0]}-{ck['band'][-1]})",
|
||||||
flush=True)
|
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(
|
adapter = MergeAdapter(
|
||||||
d=model.config.get_text_config().hidden_size).cuda()
|
d=model.config.get_text_config().hidden_size).cuda()
|
||||||
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
|
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
|
||||||
@@ -101,7 +127,7 @@ def main():
|
|||||||
max_new_tokens=args.max_new,
|
max_new_tokens=args.max_new,
|
||||||
feedforward=args.feedforward,
|
feedforward=args.feedforward,
|
||||||
inner_iters=args.inner_iters,
|
inner_iters=args.inner_iters,
|
||||||
kvmem=kvmem)
|
kvmem=kvmem, symchain=symchain)
|
||||||
if kvmem is not None:
|
if kvmem is not None:
|
||||||
from kv_memory import arm_memory
|
from kv_memory import arm_memory
|
||||||
arm_memory(None)
|
arm_memory(None)
|
||||||
|
|||||||
@@ -0,0 +1,15 @@
|
|||||||
|
# gpuq-in: results-loop/star_data.json results-loop/gsm_cot_data.json results-loop/adapter_carrycot_b1_ii10_tjs50_00_p0_e200.pt
|
||||||
|
# gpuq-out: results-loop/eval_gsm_carrycot_b1_sc*.json results-loop/train_carrycot_b1_lt03_ii10_sc*_log.json results-loop/symproj_carrycot_b1_lt03_ii10_sc*_e200.pt
|
||||||
|
git pull origin main -q 2>/dev/null
|
||||||
|
P=/home/nils/jspace/.venv/bin/python
|
||||||
|
export JLENS_MODEL=google/gemma-4-E2B-it LOOP_OUT=/home/nils/jspace/results-loop
|
||||||
|
cd /home/nils/jspace/scripts
|
||||||
|
A=$LOOP_OUT/adapter_carrycot_b1_ii10_tjs50_00_p0_e200.pt
|
||||||
|
for MODE in tf st; do
|
||||||
|
$P train_carry_cot.py --drop-steps 1 --pause-per-step 0 --base-pauses 0 \
|
||||||
|
--inner-iters 10 --lensteach 0.3 --symchain $MODE --freeze-merge \
|
||||||
|
--warm-start $A --steps 200 --lr 1e-3
|
||||||
|
$P eval_carry_cot.py --adapter $A \
|
||||||
|
--symchain $LOOP_OUT/symproj_carrycot_b1_lt03_ii10_sc${MODE}_fm_p0_e200.pt \
|
||||||
|
--tag gsm_carrycot_b1_sc$MODE --grid 0:0,2:0 --n 256 --inner-iters 10
|
||||||
|
done
|
||||||
@@ -96,6 +96,12 @@ ap.add_argument("--kvmem", type=int, default=0, metavar="CODE",
|
|||||||
ap.add_argument("--freeze-merge", action="store_true",
|
ap.add_argument("--freeze-merge", action="store_true",
|
||||||
help="freeze the (warm-started) merge adapter; train only "
|
help="freeze the (warm-started) merge adapter; train only "
|
||||||
"the kvmem adapter")
|
"the kvmem adapter")
|
||||||
|
ap.add_argument("--symchain", choices=("tf", "st"), default=None,
|
||||||
|
help="item 32: discrete latent chain — each burst tick "
|
||||||
|
"feeds back the lens-snapped token embedding through "
|
||||||
|
"a zero-init projector. tf: teacher-forced symbols "
|
||||||
|
"(ground-truth deleted-step tokens); st: free-running "
|
||||||
|
"straight-through snaps")
|
||||||
ap.add_argument("--teachstate", type=float, default=0.0, metavar="LAMBDA",
|
ap.add_argument("--teachstate", type=float, default=0.0, metavar="LAMBDA",
|
||||||
help="item 28 (Nils's variant): teacher-state distillation "
|
help="item 28 (Nils's variant): teacher-state distillation "
|
||||||
"— frozen warm-start adapter runs the FULL cot (step "
|
"— frozen warm-start adapter runs the FULL cot (step "
|
||||||
@@ -117,6 +123,7 @@ TAG = ("carrycot_ff" if ARGS.feedforward else "carrycot") + (
|
|||||||
f"_{str(ARGS.traj_fr).replace('.', '')}"
|
f"_{str(ARGS.traj_fr).replace('.', '')}"
|
||||||
if (ARGS.traj_tf or ARGS.traj_fr) else "") + (
|
if (ARGS.traj_tf or ARGS.traj_fr) else "") + (
|
||||||
f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + (
|
f"_kvm{ARGS.kvmem}" if ARGS.kvmem else "") + (
|
||||||
|
f"_sc{ARGS.symchain}" if ARGS.symchain else "") + (
|
||||||
"_fm" if ARGS.freeze_merge else "") + (
|
"_fm" if ARGS.freeze_merge else "") + (
|
||||||
f"_p{ARGS.base_pauses}" if ARGS.base_pauses >= 0 else "") + (
|
f"_p{ARGS.base_pauses}" if ARGS.base_pauses >= 0 else "") + (
|
||||||
f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + (
|
f"_ln{ARGS.lensnoise.replace(',', '_')}" if ARGS.lensnoise else "") + (
|
||||||
@@ -197,6 +204,33 @@ def gen_staging_targets(tok, cot):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def sym_burst_forward(looper, adapter, proj, ids, msk, plens, m, sym_tf,
|
||||||
|
k, lens_fn, embed_w, start_id):
|
||||||
|
"""Item 32 forward: discrete latent chain at the anchor (teacher-forced
|
||||||
|
or straight-through symbols), then the visible-token carry scan."""
|
||||||
|
from carry_common import (prompt_prefill, build_step_updates,
|
||||||
|
carry_steps, sym_iterate)
|
||||||
|
dev = ids.device
|
||||||
|
calls, _ = looper.capture(ids, msk, logits_to_keep=1)
|
||||||
|
e = looper._hin[looper.l0].detach()
|
||||||
|
B = ids.shape[0]
|
||||||
|
ar = torch.arange(ids.shape[1], device=dev)
|
||||||
|
pmask = ar[None, :] < plens[:, None].to(dev)
|
||||||
|
S, X = prompt_prefill(looper, adapter, e, calls, pmask, k)
|
||||||
|
rows = torch.arange(B, device=dev)
|
||||||
|
anchor = (plens - 1).to(dev)
|
||||||
|
states = []
|
||||||
|
S, X = sym_iterate(looper, adapter, proj, e, calls, S, X, rows, anchor,
|
||||||
|
m, lens_fn, embed_w, sym_tf=sym_tf,
|
||||||
|
start_id=start_id, use_checkpoint=True,
|
||||||
|
iter_states=states)
|
||||||
|
total = msk.sum(-1)
|
||||||
|
updates = build_step_updates(plens.to(dev), total.to(dev), dev)
|
||||||
|
S, X = carry_steps(looper, adapter, e, calls, S, X, updates,
|
||||||
|
use_checkpoint=True)
|
||||||
|
return looper.suffix_logits(S, calls), states
|
||||||
|
|
||||||
|
|
||||||
def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k,
|
def traj_burst_forward(looper, adapter, ids, msk, plens, m, T, tf_on, k,
|
||||||
mem_adapter=None):
|
mem_adapter=None):
|
||||||
"""Item 29 forward: optional teacher-forced transition predictions,
|
"""Item 29 forward: optional teacher-forced transition predictions,
|
||||||
@@ -438,6 +472,18 @@ def main():
|
|||||||
print(f"trajectory waypoints: {len(data)} items x {m} states "
|
print(f"trajectory waypoints: {len(data)} items x {m} states "
|
||||||
f"({ARGS.traj_span} span, {time.time()-t0_:.0f}s; frozen "
|
f"({ARGS.traj_span} span, {time.time()-t0_:.0f}s; frozen "
|
||||||
f"warm-start teacher, full cot)", flush=True)
|
f"warm-start teacher, full cot)", flush=True)
|
||||||
|
sym_proj, EMBED_W, START_ID = None, None, None
|
||||||
|
if ARGS.symchain:
|
||||||
|
assert ARGS.warm_start and ARGS.inner_iters and ARGS.lensteach, \
|
||||||
|
"symchain needs --warm-start, --inner-iters, --lensteach"
|
||||||
|
d_ = model.config.get_text_config().hidden_size
|
||||||
|
sym_proj = torch.nn.Linear(d_, d_).cuda()
|
||||||
|
torch.nn.init.zeros_(sym_proj.weight)
|
||||||
|
torch.nn.init.zeros_(sym_proj.bias)
|
||||||
|
EMBED_W = model.get_input_embeddings().weight.detach()
|
||||||
|
START_ID = tok("\n", add_special_tokens=False)["input_ids"][0]
|
||||||
|
print(f"symchain [{ARGS.symchain}]: zero-init proj "
|
||||||
|
f"({d_}x{d_}), start_id={START_ID}", flush=True)
|
||||||
if ARGS.teachstate:
|
if ARGS.teachstate:
|
||||||
assert ARGS.warm_start and ARGS.inner_iters, \
|
assert ARGS.warm_start and ARGS.inner_iters, \
|
||||||
"teachstate needs --warm-start (frozen teacher) + --inner-iters"
|
"teachstate needs --warm-start (frozen teacher) + --inner-iters"
|
||||||
@@ -467,6 +513,9 @@ def main():
|
|||||||
if mem_adapter is not None:
|
if mem_adapter is not None:
|
||||||
groups.append({"params": list(mem_adapter.parameters()), "lr": LR,
|
groups.append({"params": list(mem_adapter.parameters()), "lr": LR,
|
||||||
"base": LR})
|
"base": LR})
|
||||||
|
if sym_proj is not None:
|
||||||
|
groups.append({"params": list(sym_proj.parameters()), "lr": LR,
|
||||||
|
"base": LR})
|
||||||
if lora_params:
|
if lora_params:
|
||||||
groups.append({"params": lora_params, "lr": ARGS.lora_lr,
|
groups.append({"params": lora_params, "lr": ARGS.lora_lr,
|
||||||
"base": ARGS.lora_lr})
|
"base": ARGS.lora_lr})
|
||||||
@@ -498,7 +547,21 @@ def main():
|
|||||||
for g in opt.param_groups:
|
for g in opt.param_groups:
|
||||||
g["lr"] = g["base"] * lr_at(step) / LR
|
g["lr"] = g["base"] * lr_at(step) / LR
|
||||||
ckw, itstates, tfp, frs = {}, None, None, None
|
ckw, itstates, tfp, frs = {}, None, None, None
|
||||||
if ARGS.traj_tf or ARGS.traj_fr or ARGS.kvmem:
|
if ARGS.symchain:
|
||||||
|
sym_tf = None
|
||||||
|
if ARGS.symchain == "tf":
|
||||||
|
mat = torch.full((len(batch), ARGS.inner_iters), START_ID,
|
||||||
|
dtype=torch.long)
|
||||||
|
for b_i, it in enumerate(batch):
|
||||||
|
tgt = it.get("lens_targets") or []
|
||||||
|
row = [START_ID] + tgt[:-1]
|
||||||
|
mat[b_i, :len(row)] = torch.tensor(row)
|
||||||
|
sym_tf = mat.cuda()
|
||||||
|
logits, itstates = sym_burst_forward(
|
||||||
|
looper, adapter, sym_proj, ids, msk, plens,
|
||||||
|
ARGS.inner_iters, sym_tf, K_PREFILL, lens_teach, EMBED_W,
|
||||||
|
START_ID)
|
||||||
|
elif ARGS.traj_tf or ARGS.traj_fr or ARGS.kvmem:
|
||||||
T = (torch.stack([torch.as_tensor(b_["traj_states"])
|
T = (torch.stack([torch.as_tensor(b_["traj_states"])
|
||||||
for b_ in batch]).cuda()
|
for b_ in batch]).cuda()
|
||||||
if (ARGS.traj_tf or ARGS.traj_fr) else None)
|
if (ARGS.traj_tf or ARGS.traj_fr) else None)
|
||||||
@@ -608,6 +671,8 @@ def main():
|
|||||||
[p_ for p_ in adapter.parameters() if p_.requires_grad]
|
[p_ for p_ in adapter.parameters() if p_.requires_grad]
|
||||||
+ lora_params
|
+ lora_params
|
||||||
+ (list(mem_adapter.parameters()) if mem_adapter is not None
|
+ (list(mem_adapter.parameters()) if mem_adapter is not None
|
||||||
|
else [])
|
||||||
|
+ (list(sym_proj.parameters()) if sym_proj is not None
|
||||||
else []), 1.0)
|
else []), 1.0)
|
||||||
opt.step()
|
opt.step()
|
||||||
if ARGS.kvmem:
|
if ARGS.kvmem:
|
||||||
@@ -638,6 +703,9 @@ def main():
|
|||||||
if mem_adapter is not None:
|
if mem_adapter is not None:
|
||||||
torch.save(mem_adapter.state_dict(),
|
torch.save(mem_adapter.state_dict(),
|
||||||
OUT / f"kvmem_{TAG}_e{step+1}.pt")
|
OUT / f"kvmem_{TAG}_e{step+1}.pt")
|
||||||
|
if sym_proj is not None:
|
||||||
|
torch.save(sym_proj.state_dict(),
|
||||||
|
OUT / f"symproj_{TAG}_e{step+1}.pt")
|
||||||
if lora_params:
|
if lora_params:
|
||||||
torch.save({"rank": ARGS.bandlora, "band": lora_band_layers,
|
torch.save({"rank": ARGS.bandlora, "band": lora_band_layers,
|
||||||
"tensors": [p.detach().cpu()
|
"tensors": [p.detach().cpu()
|
||||||
|
|||||||
Reference in New Issue
Block a user