Skip to content

Repository files navigation

Coconut Experiments

Experiments with Meta's Coconut (Continuous Latent Space Reasoning) and vision-based code representation efficiency.

What's Here

This repo contains two related experiments:

1. Coconut: Continuous Latent Reasoning

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

2. Vision Efficiency Test

Can we represent code more efficiently as an image than as text tokens?

The experiment (a.py and low.py):

  1. Render ~60 lines of dense PyTorch code (a CausalSelfAttention block) as a PNG image
  2. 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.

Project Structure

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

Quick Start

Coconut Training

# 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

Vision Efficiency Test

# 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

Training Modes

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

Key Parameters

  • c_thought: Number of continuous thought tokens per reasoning step
  • epochs_per_stage: Training epochs before adding more latent tokens
  • max_latent_stage: Maximum number of reasoning steps to replace with latents
  • pad_latent_to_max: Pad shorter sequences to max latent count

Datasets

  • GSM8K: Grade school math problems
  • ProntoQA: Logical reasoning (5-hop)
  • ProsQA: Procedural reasoning (included in data/)

Requirements

  • 4× A100 80GB GPUs (for full training)
  • PyTorch 2.0+
  • Transformers
  • wandb account for logging

References

License

MIT License (see LICENSE)

About

Continuous-latent reasoning and vision-as-context experiments inspired by Coconut

Topics

Resources

Code of conduct

Contributing

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages