What Matters for Latent Reasoning with Flow Matching
Organizations: Samsung AI Cambridge · Technical University of Iasi · Queen Mary University of London
Abstract
Latent reasoning lets a large language model (LLM) think in a continuous space and verbalize only the answer. We argue that an effective latent thought must meet five requirements: it should be useful, helping produce the correct answer rather than merely changing it, diverse, so that resampling yields different reasoning trajectories, explainable, so that a decoded chain of thought (CoT) reflects reasoning the answer actually follows, refinable with more inference compute, and efficient, costing less than an explicit CoT at comparable accuracy. Current methods rarely meet these requirements: they learn shortcuts from the question, distill the explicit CoT into their weights, or imitate it one token at a time. We focus on flow matching in a learned latent space, the family we argue is best placed to meet them, and identify the training choices that make it work. The result is Flow-based Latent Reasoning (FLaRe), a simple recipe covering what the latent space encodes and how to shape it, where to train the flow, how to read out the answer, and a final stage of training on the model's own verified thoughts. A probe for each requirement shows that FLaRe improves on prior latent methods in all five. It also compares favorably with them on arithmetic benchmarks, while reaching 97% of the accuracy of explicit CoT at a quarter of its latency.
Figures & tables
| Part of the loop | Variant | Direct | Decoded |
|---|---|---|---|
| Stage 1 | 55.3 | 61.3 | |
| Stage 2 | 59.1 | 62.6 | |
| Training | No online flow loss | 59.0 | 61.6 |
| No stage 1 replay | 59.0 | 62.5 | |
| Targets | No verification | 57.1 | 60.7 |
| One trace | 58.8 | 61.8 | |
| Useful | Diverse | Explainable | Refinable | Efficient | ||
| Method | follows | pass@16 | vote | values in order | CoT compute | speedup at % CoT acc. |
| Explicit CoT | 97.2 | – | – | – | – | at 100% |
| Coconut | 7.3 | 19.1 | at 61% | |||
| CODI | 31.8 | 38.5 | at 94% | |||
| PCCoT | 27.7 | 7.8 | at 91% | |||
| FLaRe, stage 1 | 39.7 | 46.0 | / | at 93% | ||
| Qwen2.5-0.5B-Instruct | Llama-3.2-1B-Instruct | Llama-3.2-3B-Instruct | ||||||||||
| IID | OOD | IID | OOD | IID | OOD | |||||||
| Method | GSM8K | Hard | SVAMP | MArith | GSM8K | Hard | SVAMP | MArith | GSM8K | Hard | SVAMP | MArith |
| Explicit CoT | 59.7 | 17.0 | 61.6 | 98.3 | 62.5 | 14.7 | 66.5 | 97.8 | 72.4 | 20.8 | 73.4 | 98.3 |
| No-CoT | 27.3 | 6.4 | 38.5 | 57.8 | 35.0 | 7.4 | 38.1 | 67.8 | 39.2 | 9.3 | 56.7 | 96.7 |
| Coconut | 28.2 | 6.8 | 41.9 | 58.3 | 36.1 | 8.3 | 48.4 | 87.2 | 45.3 | 10.9 | 57.5 | 96.7 |
| iCoT | 32.9 | 7.4 | 41.0 | 66.7 | 36.5 | 8.4 | 40.8 | 72.8 | 42.1 | 9.6 | 50.7 | 91.1 |
| Method | Family | Thought lives in | Produced by | Trained with | Decodable | Budget |
|---|---|---|---|---|---|---|
| Explicit CoT | Horizontal | Tokens | AR decoding | CoT SFT | Written out | Model-chosen |
| iCoT | Vertical | Hidden states | Single pass | CoT curriculum | No | Fixed |
| Pause tokens | Vertical | Hidden states | Single pass | Answer loss | No | Fixed |
| Coconut | Horizontal | Embeddings | AR feedback | CoT curriculum | LM head | Fixed |
| CODI | Horizontal | Embeddings | AR feedback | CoT distillation | LM head | Fixed |
| PCCoT | Parallel | Embeddings | Jacobi iteration | CoT distillation | LM head | Fixed |
Appendix figures & tables13 assets
Supplementary material from the paper’s appendix.
Appendix
| VAE | |||
| The VAE is initialized from a model trained for one epoch on the natural language CoTs of OMI-2 ( Toshniwal et al., 2025 ) with 32 slots and a single natural language route. It is then reduced to slots at initialization and trained with the dual objective of equation 3 with both terms weighted equally. Token substitution draws the replacement uniformly from the vocabulary with clean decoder targets, and the latent noise and dropout act on the sampled code before decoding. The statistics are the mean and standard deviation of the posterior means of 20K training rows. | |||
| Backbones , | Llama-3.2-1B / 3B, fine-tuned | Token substitution | 0.3 on the encoder input |
| Initialization | OMI-2 language VAE, | Latent noise | VP, , on half of the codes |
| Latent code | Latent dropout | 0.4 | |
| Input, decoder routes | , dual routes, equal weights | Optimizer | AdamW, , wd , clip 1.0 |
| Training data | paired GSM8K-Aug, 367K rows | Learning rate | encoder, decoder |
| Method | Qwen2.5-0.5B-Instruct | Llama-3.2-1B-Instruct | Llama-3.2-3B-Instruct |
|---|---|---|---|
| Explicit CoT | T | T | R (a) |
| No-CoT | T | T | T |
| iCoT | T | T | T |
| Coconut | T | R (b) | T |
| CODI | T | P | T |
| PCCoT | T | P on GSM8K, R (c) on the OOD sets | T |
| Method | Tuning | Learning rate | Epochs | Other settings |
|---|---|---|---|---|
| Explicit CoT | LoRA, , | / / – | 10 | batch 128, cosine, 3% warmup, wd 0.1 |
| No-CoT | LoRA, , | / / | 10 | as explicit CoT, answer only |
| iCoT | full, single precision | 20 | batch 32, 8 CoT tokens removed per epoch (Qwen: 11) | |
| Coconut | full, bf16 | 3 + 6 + 1 | , a CoT stage, stages 1 to 6 and a fully latent stage, batch 128 | |
| CODI | LoRA, , | / – / | 10 / – / 8 | 6 latents, distillation weight 20, batch 128 |
| PCCoT | LoRA, , | / – / | 10 | 24 latents, iterations, batch 128 |
| Corpus | Origin | Rows | Questions |
|---|---|---|---|
| GSM8K-Aug, flow | whynlp/gsm8k-aug train | 385K | 385K |
| GSM8K-Aug, paired | aligned with whynlp/gsm8k-aug-nl | 367K | 367K |
| Diverse CoT | Qwen2.5-32B-Instruct on GSM8K-Aug questions | 216K | 148K |
| Diverse CoT, paired | same, with a verified natural language twin | 205K | 142K |
| IID data | OMI-2 augmented_gsm8k , converted | 166K | 39K |
| OOD data | OMI-2 augmented_math , converted | 463K | 80K |
| VAE training data | Recon. | Direct | Decoded |
|---|---|---|---|
| GSM8K-Aug CoT | 99.1 | 50.7 | 56.9 |
| + Diverse CoT | 99.0 | 48.8 | 54.8 |
| + IID data | 99.0 | 51.6 | 56.9 |
| + OOD data | 99.3 | 50.9 | 56.0 |
| + IID and OOD data ( budget) | 99.6 | 51.1 | 58.3 |
| IID pair, from the augmented_gsm8k share | |
|---|---|
| Question | Lucy has orange trees that produce 4 oranges each. If an orange can be sold for 80, how many orange trees does she need to harvest? |
| Each orange tree produces 4 oranges. Each orange can be sold for 2 = 80, then she needs to harvest 8 = 10 trees. Thus, Lucy needs to harvest 10 orange trees. | |
| << 4*2=8 >> << 80/8=10 >> << 10 >> | |
| 10 | |
| OOD pair, from the augmented_math share | |
| Question | The greatest common divisor of , , and is . What is the smallest possible value of that is greater than ? |
| Model | Follows | Keeps original | Neither | Clean follow |
|---|---|---|---|---|
| Explicit CoT (reference steps) | 97.2 | 0.2 | 2.5 | 99.9 |
| CODI | 31.8 | 19.8 | 48.4 | 68.4 |
| PCCoT | 27.7 | 20.8 | 51.6 | 62.3 |
| Coconut | 7.3 | 29.5 | 63.2 | 19.2 |
| FLaRe, stage 1 | 39.7 | 8.7 | 51.6 | 91.4 |
| FLaRe, stage 2 | 37.3 | 15.0 | 47.7 | 76.1 |
| Model | Greedy | pass@1 | pass@ | Majority | Distinct ans. | Committed err. | |
|---|---|---|---|---|---|---|---|
| CODI | 0.5 | 55.6 | 54.6 | 59.1 | 54.9 | 1.46 | 51% |
| PCCoT | 1.0 | 54.1 | 51.9 | 54.9 | 52.2 | 1.31 | 58% |
| Coconut | 2.0 | 36.0 | 33.8 | 40.6 | 34.5 | 1.89 | 34% |
| FLaRe, stage 1, decoded | – | 61.2 | 60.1 | 77.8 | 64.4 | 3.41 | 6% |
| FLaRe, stage 1, direct | – | 55.3 | – | 67.1 | 55.6 | 2.68 | – |
| FLaRe, stage 2, decoded | – | 62.6 | 62.0 | 77.3 | 64.4 | 2.89 | 9% |