Reproduction of the 2026 workspace/J-lens paper on gemma-4 (E2B/12B/26B), plus the workspace-loop retrofit line: merge adapter, prompt-only latent planning (MBPP), carry variant, attribution controls (FF/pause/untrained), band-location ablation, Blocksworld harness, 12B replication scripts. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
106 lines
4.0 KiB
Python
106 lines
4.0 KiB
Python
"""Loop convergence dynamics: does the recurrence reach a fixed point?
|
||
|
||
For trained vs untrained merge, trace across iterations k:
|
||
- cos(s_k, s_{k-1}) (mean over positions) -> fixed point if -> 1
|
||
- |s_k| / |e| -> norm control
|
||
- P('spider') under the J-lens at L30 -> what the state converges TO
|
||
Averaged over the spider prompt + a few GSM8K test questions.
|
||
"""
|
||
|
||
import json
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
import matplotlib
|
||
matplotlib.use("Agg")
|
||
import matplotlib.pyplot as plt
|
||
import torch
|
||
|
||
from loop_common import BandLooper, MergeAdapter, chat_prompt, DIRECT_SUFFIX
|
||
|
||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||
from jlens.core import load_model # noqa: E402
|
||
|
||
OUT = Path(__file__).resolve().parent.parent / "results-loop"
|
||
KMAX = 10
|
||
BLUE, GRAY = "#2b6cb0", "#8a8f98"
|
||
|
||
|
||
@torch.no_grad()
|
||
def trace(looper, adapter, tok, prompt_text):
|
||
ids = tok(prompt_text, return_tensors="pt",
|
||
add_special_tokens=False)["input_ids"].cuda()
|
||
calls, _ = looper.capture(ids)
|
||
e = looper._hin[looper.l0]
|
||
s = looper.band(e, calls)
|
||
rows = []
|
||
for k in range(1, KMAX + 1):
|
||
new = looper.band(adapter(e, s), calls)
|
||
cos = torch.nn.functional.cosine_similarity(
|
||
new[0].float(), s[0].float(), dim=-1).mean().item()
|
||
rows.append({"k": k, "cos": cos,
|
||
"norm": (new.norm() / e.norm()).item()})
|
||
s = new
|
||
return rows
|
||
|
||
|
||
def main():
|
||
model, tok = load_model(dtype=torch.bfloat16)
|
||
looper = BandLooper(model)
|
||
|
||
gsm = [it for it in json.load(open(OUT / "star_data.json"))
|
||
if it["split"] == "test"][:3]
|
||
prompts = [chat_prompt(tok, "The animal that spins webs has how many legs? "
|
||
"Answer with just the number.", "")]
|
||
prompts += [chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in gsm]
|
||
|
||
curves = {}
|
||
for tag, path in (("untrained", None), ("trained", OUT / "adapter.pt")):
|
||
adapter = MergeAdapter().cuda()
|
||
if path:
|
||
adapter.load_state_dict(torch.load(path, map_location="cuda"))
|
||
traces = [trace(looper, adapter, tok, p) for p in prompts]
|
||
curves[tag] = {
|
||
"cos": [sum(t[i]["cos"] for t in traces) / len(traces)
|
||
for i in range(KMAX)],
|
||
"norm": [sum(t[i]["norm"] for t in traces) / len(traces)
|
||
for i in range(KMAX)],
|
||
}
|
||
print(tag, "cos:", [round(c, 3) for c in curves[tag]["cos"]], flush=True)
|
||
print(tag, "norm:", [round(c, 3) for c in curves[tag]["norm"]], flush=True)
|
||
|
||
ks = list(range(1, KMAX + 1))
|
||
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4))
|
||
for ax in (ax1, ax2):
|
||
ax.grid(True, color="#e5e5e5", lw=0.7)
|
||
ax.set_axisbelow(True)
|
||
for sp in ("top", "right"):
|
||
ax.spines[sp].set_visible(False)
|
||
ax.set_xticks(ks)
|
||
ax.set_xlabel("loop iteration k")
|
||
ax.axvspan(2, 4, color="#f2e8cf", alpha=0.45, zorder=0)
|
||
|
||
for tag, c in (("trained", BLUE), ("untrained", GRAY)):
|
||
ax1.plot(ks, curves[tag]["cos"], "-o", color=c, lw=2, ms=5, label=tag)
|
||
ax2.plot(ks, curves[tag]["norm"], "-o", color=c, lw=2, ms=5, label=tag)
|
||
ax1.set_ylabel("cos(s_k, s_{k−1}) (mean over positions)")
|
||
ax1.set_title("Successive-state similarity: fixed point?", fontsize=11)
|
||
ax1.legend(fontsize=8, frameon=False, loc="lower right")
|
||
ax1.text(3, ax1.get_ylim()[0] + 0.02 * (ax1.get_ylim()[1] - ax1.get_ylim()[0]),
|
||
"accuracy &\nsharpening plateau", fontsize=7.5, color="#7a5a00",
|
||
ha="center")
|
||
ax2.set_ylabel("|s_k| / |e|")
|
||
ax2.set_title("State norm across iterations", fontsize=11)
|
||
ax2.legend(fontsize=8, frameon=False)
|
||
|
||
fig.suptitle("Loop dynamics (spider + 3 GSM8K prompts, mean)", fontsize=12,
|
||
y=1.02)
|
||
fig.tight_layout()
|
||
fig.savefig(OUT / "loop_dynamics.png", dpi=140, bbox_inches="tight",
|
||
facecolor="white")
|
||
print("wrote", OUT / "loop_dynamics.png")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|