erythropygia/qwen3-4b-latent-reasoning-retrofit
Latent Reasoning Retrofit — LoRA for Qwen3-4B-Thinking-2507
This is an attempt to build an Astra-style hidden thinking phase on an open 4B checkpoint, and to measure how far that attempt actually gets.
The structure described for closed frontier models is a think phase that runs as hidden state rather than emitted tokens: fewer output tokens, and a chain of thought that can no longer be read off the transcript. The mechanism is straightforward to implement. At the <think> position you stop decoding tokens and instead feed the final hidden state back in as an input embedding, repeatedly. Each pass writes a KV slot that the answer can attend to and carries no token ID, so nothing about it appears in the output.
That part works. The reasoning does not disappear with it.
Force </think> after the latent phase and the model writes its derivation immediately afterwards — same wording, same steps, on the other side of the tag. And a control that runs zero latent steps scores the same as 32 of them, so the latent steps are not what produces the answer. What they do measurably change is length: with the layer loop engaged, each latent step removes roughly 20 written tokens.
So: the depth mechanism is real and helps accuracy, the substitution mechanism is real and shortens output, and the property the whole exercise was aimed at — reasoning that stops being visible — does not appear at this budget.
Try it on your own input
test_live.py loads the model once and then loops, so you type a question and watch the answer stream out in every mode. Nothing is truncated.
python test_live.py
python test_live.py --prompt "A robe takes 2 bolts of blue fiber and half that much white fiber. How many bolts in total?"
python test_live.py --modes readable,skip,latent32 --loopWhat the model actually receives
Your question goes through Qwen3's thinking chat template, which ends the prompt with an open <think> tag. That last position is where all the modes diverge:
<|im_start|>user
A robe takes 2 bolts of blue fiber and half that much white fiber. How many bolts in total?<|im_end|>
<|im_start|>assistant
<think>- readable decodes tokens from there, so the think block gets written out.
- skip injects
</think>immediately and decodes the answer. No latent steps at all. - latentN runs N hidden-state feedback passes first, then injects
</think>.
What comes back
latent_generate returns a dict; text is the full decoded output and answer_tokens counts it. Real output from the run above, --modes skip,latent32 --max-new 400:
OUTPUT -- skip 400 tokens, <think> block 0
</think>
First, the blue fiber part is straightforward: 2 bolts. Then the white fiber is half that
much. Wait, "half that much" -- that "that much" probably refers to the blue fiber. So half
of 2 bolts would be 1 bolt. Let me make sure. [...] Maybe someone could misinterpret "half
that much" as half the total, but the wording is "half that much white fiber" [...]
(still going when the cap hit)
OUTPUT -- latent32 260 tokens, <think> block 0
</think>
The robe takes 2 bolts of blue fiber and 1 bolt of white fiber. Adding those together, the
total number of bolts needed is 2 + 1 = 3.
So, the answer should be 3 bolts in total.
Wait, let me make sure I didn't misread the question. [...] I don't think I made a mistake
here. The answer should be 3 bolts total.Both rows report <think> block: 0, and both write several hundred tokens of reasoning anyway. That is the finding in one screen: the block is gone, the reasoning is not. The latent run is shorter (260 against 400+) and commits to an answer sooner, which is the part that does work.
The same thing measured on one GSM8K problem
GSM8K test problem 0 ("Janet's ducks lay 16 eggs per day...", answer 18), run through every condition with python test.py --gsm8k 0 --r 8 --r 32 --max-new 2048 --lens. All four conditions answer 18.
The <think> column is where the naive reading goes wrong. Three of these rows write no think block at all, and one of them still emits 1150 tokens.
readable — the block the model normally writes:
Okay, let's try to figure out how much Janet makes every day at the farmers' market. Hmm, first, let's break down the problem step by step. [...] Wait, does that mean three eggs per day? Because it says "every morning" [...]
skip — no latent steps, </think> injected at the fork point:
</think> Okay, let's try to figure out how much Janet makes every day at the farmers' market. Hmm, first, let's break down the problem step by step. [...] Wait, does that mean three eggs each morning? The problem says "three for breakfast every morning" [...]Same opening sentence, same hedging, same derivation — on the other side of the tag. This is the whole result in one comparison. Nothing was hidden; the text moved.
latent, 32 steps — the substitution actually doing something:
</think> First, let's break down the problem step by step. Janet's ducks lay 16 eggs each day. Here's what happens to those eggs: 1. Eggs eaten for breakfast: Janet eats 3 eggs every morning. [...]573 tokens against 1150 for the control, and the character of the text changes too: the exploratory "Wait, does that mean..." passes are gone and what is left is a clean write-up. That is what the latent steps buy — a shorter, more committed answer, not a hidden one.
What the latent vectors look like (--lens, nearest token embeddings):
prefill h 'Okay':0.144 'This':0.137 ' This':0.123
step 3 'The':0.104 'First':0.103 'Wait':0.101
step 4 '</think>':0.121 'The':0.095 ' the':0.092
step 5 'Wait':0.095 'First':0.086 'So':0.082
step 11 'Wait':0.101 'The':0.096 ' the':0.093The nearest neighbours are the discourse markers a chain of thought is built from — Wait, First, So, </think>. Suggestive, but the cosines are ~0.10, which is not a match; these vectors are not in token space. That is the one sense in which the latent phase is opaque, and it is worth nothing here, because the reasoning is printed in full a moment later.
What was tried
Two independent mechanisms, measured separately because they are separately true:
- Axis A — latent depth. A window of decoder layers runs K times with damped Euler sub-steps, so each token is computed through a deeper effective stack. No weights change and
K=1is bit-identical to the base model. Needs no training. - Axis B — latent substitution. The hidden-state feedback described above. This is what the adapter trains, and it is the mechanism the interesting claim rests on.
Axis A does not produce axis B. That was measured before training: the loop wrapper moved accuracy but left every substitution contrast null (p ≥ 0.18). Hence the curriculum.
Results
Base: Qwen/Qwen3-4B-Thinking-2507 · Loop window: [15-18] · K: 2 · Strategy: Euler · Eval: GSM8K test, greedy, chat template applied, paired McNemar.
Axis A: latent depth works without training
The loop wrapper alone: 86.0% → 91.0% on GSM8K (+5.0 pp, n=200, p=0.021), measured before any training, so the wrapper is the only thing that changed.
Axis B: dependence on the written trace drops
Damage the model's own reasoning trace and see whether the answer survives. A trace that is carrying the reasoning degrades as you cut it; one that has stopped carrying it does not.
Paired against base: truncate_50 +23.0 pp (p<0.001), truncate_25 +18.0 pp (p<0.001), corrupt +13.0 pp (p=0.001). Deleting the second half of its own trace costs the base model 27 points and costs this adapter nothing. Corrupting a single digit costs the base model 27 points and costs this adapter 9.
Its traces are also 31% longer, so the same percentage is more text. In absolute tokens the effect survives: this adapter scores 73.0% on 286 trace tokens where the base model scores 58.5% on 436.
The latent steps are not what makes it work
The control forces </think> with zero latent steps and lets the model answer.
n=60, so the standard error is about 4.3 pp and every row here sits inside one band, control included. This experiment does not establish an accuracy benefit from the latent steps.
Output length is a different story: it falls monotonically with the number of latent steps, about 20 written tokens per step with the loop engaged, and roughly 24 per step at small counts — which is the exchange rate the curriculum was built around. Total compute lands at 0.71–0.76× the written baseline at comparable accuracy.
Logit-lens cosine between the latent vectors and their nearest token embeddings stays around 0.10, so those vectors are not sitting in ordinary token space. It buys no privacy, because the derivation is written out in plain text a moment later.
What this is not
- Not a reproduction of a frontier latent-reasoning system. The budget is roughly 5,000× smaller — 150M LoRA tokens against recurrent pretraining at scale.
- Not a model that hides its reasoning. It writes reasoning in text. If you need an uninspectable think phase, this is not that, and the tables above are the evidence.
- Not an accuracy win. On full GSM8K it scores 81.0% against 86.0% for the base model and 91.0% for the untrained loop. See limitation 1 before reading that as damage.
- Not a clean decomposition. 65% of training steps used ordinary LM loss and no loop-free LoRA ablation was run, so "loop + healing" cannot be separated from plain SFT.
- Not a learned habit. What is trained is a switch, not a policy: the latent path only does anything when you turn it on. Normal generation is unchanged.
Limitations
- The accuracy drop is unresolved. This adapter's traces are 31% longer and hit the 2048-token generation cap far more often (16 of 17 batches, against 13 for the base model), so a reasoning failure and a truncated trace look identical in that number. A rerun at 4096 would separate them. It has not been done.
- The truncation battery uses two different answer budgets.
fullgenerates up to 2048 tokens; every damaged condition gets 256. Sotruncate_0 ≈ 10%means "trace removed and no room to redo it" — with 2048 tokens and no trace at all this checkpoint still scores 85%. Cross-arm comparisons use one protocol and stay valid; the absolute numbers do not mean what they look like. - The latent chain is numerically unstable. Token decoding has an argmax that re-quantises noise every step; the latent chain has no equivalent. A change of batch shape moves the prefill by cosine 0.999, which the chain opens to ~0.5 within 8 steps. This is not a padding bug — identical rows in one batch are bit-identical. Aggregate accuracy is a valid estimate; which problems come out right is not stable across batch sizes.
- Limited statistical power.
n=200for the truncation battery,n=60for the substitution table. The large truncation effects clear the noise; the accuracy differences do not. - One task. GSM8K only.
Usage
The adapter loads as an ordinary PEFT LoRA. The latent path additionally needs the loop wrapper and latent_projector.pt, both of which ship in this repo.
The projector must stay in fp32. What it learned is a small perturbation of the identity — off-diagonal weights around 1.5e-3 — and bf16 quantises near 1.0 in steps of about 7.8e-3, which erases it.
import torch
from huggingface_hub import hf_hub_download
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer
from looped import LatentProjector, LatentThinkConfig, LoopConfig, apply_loop, latent_generate
from looped.latent_think import embedding_rms
REPO = "erythropygia/qwen3-4b-latent-reasoning-retrofit"
BASE = "Qwen/Qwen3-4B-Thinking-2507"
model = AutoModelForCausalLM.from_pretrained(BASE, dtype=torch.bfloat16, device_map="cuda:0")
tok = AutoTokenizer.from_pretrained(BASE)
model = PeftModel.from_pretrained(model, REPO).eval()
proj = LatentProjector(model.config.hidden_size, "linear",
embedding_rms(model)).to(model.device, torch.float32)
proj.load_state_dict(torch.load(hf_hub_download(REPO, "latent_projector.pt"),
map_location="cpu"))
# Axis A, optional. Leave it out to run latent substitution on its own.
apply_loop(model, LoopConfig(start=15, end=18, K=2, strategy="euler",
cache="first", mode="block", decode_mode="full"))
ids = tok(tok.apply_chat_template([{"role": "user", "content": "..."}],
tokenize=False, add_generation_prompt=True),
return_tensors="pt").input_ids.cuda()
out = latent_generate(model, tok, ids,
LatentThinkConfig(r_latent=32, projector="linear"), projector=proj)
print(out["text"][0])Two scripts ship with this repo:
test_live.py— interactive. Type a question, watch every mode stream its full answer.--all(or/allin session) runs every value the adapter was trained at: K 1 through 8, and 4 through 32 latent steps, against the readable and zero-latent references./modes latent32@K2,/tokens Nchange what runs without reloading the model.test.py— one prompt through every condition in one shot, plus the logit-lens readout of each latent vector (--lens) and GSM8K problems by index (--gsm8k N).
Both run the zero-latent-step control alongside the latent modes. Run it before drawing conclusions from any latent output.
On a 3090 decoding runs at about 15 tokens/second at batch size 1, and that is a launch-latency floor rather than a compute limit: step time is essentially identical at batch 1 and batch 12 (69 ms against 71 ms). Both scripts merge the adapter into the base weights at load time, which is worth 35% — an unmerged LoRA costs 102 ms per token against 67 ms merged, because every target module adds two extra small matmuls per token. Pass --no-merge to keep it unmerged.
Training
The K=1 KL anchor is what keeps the base model's behaviour intact while the K>1 steps learn to use the extra depth.
Citation
The mechanisms come from prior work and are not introduced here:
- Coconut — latent/continuous thought (Hao et al., 2024)
- Huginn — truncated backpropagation through a recurrent depth chain (Geiping et al., 2025)
- the depth-recurrence literature — the layer-loop window
What this repository adds is the measurement protocol and a negative result at 4B scale.
Framework versions
- PEFT 0.20.0
