Experiments with Meta's Coconut (Continuous Latent Space Reasoning) and vision-based code representation efficiency.
This repo contains two related experiments:
Based on the paper "Training Large Language Models to Reason in a Continuous Latent Space" by Meta AI.
The core idea: Instead of generating chain-of-thought reasoning as text tokens, train models to reason using continuous hidden state vectors (latent thoughts).
How it works:
- Special tokens
<|start-latent|>,<|latent|>,<|end-latent|>mark reasoning sections - When the model hits a
<latent>token, instead of using its embedding, it feeds back the hidden state from the previous position - This creates a "thought loop" in continuous embedding space rather than discrete token space
- The model learns to compress multi-step reasoning into these latent representations
Benefits:
- More compute-efficient (no text generation for reasoning steps)
- Potentially more expressive (not constrained to natural language)
- Internal reasoning is not exposed in output
Can we represent code more efficiently as an image than as text tokens?
The experiment (a.py and low.py):
- Render ~60 lines of dense PyTorch code (a CausalSelfAttention block) as a PNG image
- Downscale to different resolutions to test readability vs token efficiency:
- 512×512px (~1300 ViT patches) - "Half-res" mode
- 384×384px (~750 patches) - Native ViT resolution, break-even with text
- 256×256px (~330 patches) - Extreme compression, more efficient than text if readable
The question: At what resolution can a vision model still parse the code? If 256px works, visual code representation is strictly more token-efficient than text.
coconut/
├── coconut.py # Core Coconut model (latent reasoning loop)
├── run.py # Distributed training script (FSDP + wandb)
├── dataset.py # Data loading with <latent> token handling
├── utils.py # Config and seed utilities
├── a.py # Generate code-as-image artifact
├── low.py # Downscale image for efficiency testing
├── args/ # Training configs (GSM8K, ProntoQA, ProsQA)
├── data/ # ProsQA dataset
├── preprocessing/ # Dataset preparation scripts
└── assets/ # Coconut diagram
# Setup
conda create --name coconut python=3.12
conda activate coconut
pip install -r requirements.txt
# Prepare GSM8K data
bash preprocessing/gsm_icot.bash
# Train CoT baseline (stage 0)
torchrun --nnodes 1 --nproc_per_node 4 run.py args/gsm_cot.yaml
# Train Coconut (continuous latent reasoning)
torchrun --nnodes 1 --nproc_per_node 4 run.py args/gsm_coconut.yaml# Generate high-res code image
python a.py
# Create downscaled versions for testing
python low.py
# Output: eff_test_512px.png, eff_test_384px.png, eff_test_256px.png| Mode | Description |
|---|---|
coconut: True |
Latent reasoning with continuous thoughts |
cot: True |
Standard chain-of-thought (text reasoning) |
no_thoughts: True |
Coconut architecture but 0 latent tokens |
no_cot: True |
Direct answer, no reasoning |
c_thought: Number of continuous thought tokens per reasoning stepepochs_per_stage: Training epochs before adding more latent tokensmax_latent_stage: Maximum number of reasoning steps to replace with latentspad_latent_to_max: Pad shorter sequences to max latent count
- GSM8K: Grade school math problems
- ProntoQA: Logical reasoning (5-hop)
- ProsQA: Procedural reasoning (included in
data/)
- 4× A100 80GB GPUs (for full training)
- PyTorch 2.0+
- Transformers
- wandb account for logging
MIT License (see LICENSE)