E1c: k*=0 routing + pre-loop halt target (item 19 scored, E1c pre-registered)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -408,3 +408,15 @@ points at a repair).
|
|||||||
deep); (d) overall >= 52.0 at E[k] <= 2.2. If (c) fails while (a,b)
|
deep); (d) overall >= 52.0 at E[k] <= 2.2. If (c) fails while (a,b)
|
||||||
hold, halting-head recall saturates at probe level and gate quality,
|
hold, halting-head recall saturates at probe level and gate quality,
|
||||||
not gate training, is the binding constraint.
|
not gate training, is the binding constraint.
|
||||||
|
|
||||||
|
--- Outcome, item 19 / E1b (scored 2026-07-16 ~01:15): prediction (c)
|
||||||
|
CONFIRMED (hard 39.3 >= 35.7 at mean k* 2.18), (a) FAILED at r=0.217
|
||||||
|
(selectivity real — hard routed 2x deeper than easy (2.18 vs 1.08), the
|
||||||
|
program's first nonzero gate correlation — but weak at deploy), (b,d)
|
||||||
|
FAILED for a traced design reason: halted_k_per_item lacked k*=0, so easy
|
||||||
|
items were forced through >=1 iteration and landed on the merge's WORST
|
||||||
|
easy depth (k=1: 85.2%); E0's 97.5% came precisely from k=0 routing.
|
||||||
|
E1c amendment (pre-registered before running, same session): pre-loop
|
||||||
|
halt consult on s_0 enabling k*=0; targets easy->0, hard->4; threshold
|
||||||
|
0.5 unchanged (calibration deferred unless E1c misses). Predictions:
|
||||||
|
easy >= 95%, hard >= 35.7%, r >= 0.4, overall >= 51.2 at E[k] <= 1.5.
|
||||||
|
|||||||
@@ -78,17 +78,26 @@ def halting_loop_logits(looper, adapter, input_ids, k_max, attention_mask,
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def halted_k_per_item(looper, adapter, input_ids, k_max, attention_mask,
|
def halted_k_per_item(looper, adapter, input_ids, k_max, attention_mask,
|
||||||
prompt_last_idx, thresh=0.5):
|
prompt_last_idx, thresh=0.5, allow_k0=True):
|
||||||
"""Deploy-time halting: smallest k where cumulative halt mass >= thresh
|
"""Deploy-time halting: smallest k where cumulative halt mass >= thresh
|
||||||
(k_max if never). Returns (B,) ints in 1..k_max."""
|
(k_max if never). With allow_k0, the head is also consulted on the
|
||||||
|
pre-loop state s_0 — items it flags there get k*=0 (no looping at
|
||||||
|
all; the E0 lesson: easy items are safest untouched).
|
||||||
|
Returns (B,) ints in 0..k_max."""
|
||||||
calls, _ = looper.capture(input_ids, attention_mask, logits_to_keep=1)
|
calls, _ = looper.capture(input_ids, attention_mask, logits_to_keep=1)
|
||||||
e = looper._hin[looper.l0].detach()
|
e = looper._hin[looper.l0].detach()
|
||||||
B = e.shape[0]
|
B = e.shape[0]
|
||||||
bidx = torch.arange(B, device=e.device)
|
bidx = torch.arange(B, device=e.device)
|
||||||
s = looper.band(e, calls)
|
s = looper.band(e, calls)
|
||||||
keep = torch.ones(B, device=e.device)
|
|
||||||
cum = torch.zeros(B, device=e.device)
|
|
||||||
kstar = torch.full((B,), k_max, dtype=torch.long, device=e.device)
|
kstar = torch.full((B,), k_max, dtype=torch.long, device=e.device)
|
||||||
|
cum = torch.zeros(B, device=e.device)
|
||||||
|
keep = torch.ones(B, device=e.device)
|
||||||
|
if allow_k0:
|
||||||
|
p0 = adapter.halt_prob(e[bidx, prompt_last_idx],
|
||||||
|
s[bidx, prompt_last_idx])
|
||||||
|
kstar[p0 >= thresh] = 0
|
||||||
|
cum = cum + p0
|
||||||
|
keep = keep * (1 - p0)
|
||||||
for i in range(k_max):
|
for i in range(k_max):
|
||||||
x = adapter(e, s)
|
x = adapter(e, s)
|
||||||
s = looper.band(x, calls)
|
s = looper.band(x, calls)
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ ap.add_argument("--seed", type=int, default=0)
|
|||||||
ap.add_argument("--lr", type=float, default=3e-3)
|
ap.add_argument("--lr", type=float, default=3e-3)
|
||||||
ARGS = ap.parse_args()
|
ARGS = ap.parse_args()
|
||||||
BATCH = 4
|
BATCH = 4
|
||||||
TARGET_K = {"easy": 1, "hard": ARGS.kmax}
|
TARGET_K = {"easy": 0, "hard": ARGS.kmax} # easy: halt BEFORE looping
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
@@ -81,7 +81,9 @@ def main():
|
|||||||
calls, _ = looper.capture(ids, msk, logits_to_keep=1)
|
calls, _ = looper.capture(ids, msk, logits_to_keep=1)
|
||||||
e = looper._hin[looper.l0].detach()
|
e = looper._hin[looper.l0].detach()
|
||||||
s = looper.band(e, calls)
|
s = looper.band(e, calls)
|
||||||
loss = 0.0
|
# pre-loop halt decision on s_0: target 1 for easy (k*=0)
|
||||||
|
p0 = adapter.halt_prob(e[bidx, pl], s[bidx, pl])
|
||||||
|
loss = F.binary_cross_entropy(p0.float(), (tk == 0).float())
|
||||||
for i in range(ARGS.kmax):
|
for i in range(ARGS.kmax):
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
x = adapter(e, s)
|
x = adapter(e, s)
|
||||||
@@ -90,7 +92,7 @@ def main():
|
|||||||
p = adapter.halt_prob(e[bidx, pl], s[bidx, pl])
|
p = adapter.halt_prob(e[bidx, pl], s[bidx, pl])
|
||||||
tgt = ((i + 1) >= tk).float()
|
tgt = ((i + 1) >= tk).float()
|
||||||
loss = loss + F.binary_cross_entropy(p.float(), tgt)
|
loss = loss + F.binary_cross_entropy(p.float(), tgt)
|
||||||
loss = loss / ARGS.kmax
|
loss = loss / (ARGS.kmax + 1)
|
||||||
opt.zero_grad(set_to_none=True)
|
opt.zero_grad(set_to_none=True)
|
||||||
loss.backward()
|
loss.backward()
|
||||||
opt.step()
|
opt.step()
|
||||||
|
|||||||
Reference in New Issue
Block a user