CEDR / README.md
xc91's picture
Clarify Qwen's role as representation teacher
a12277e verified
|
Raw History Blame Contribute Delete
5.46 kB
metadata
language:
  - en
license: apache-2.0
library_name: cedr
pipeline_tag: text-generation
tags:
  - continuous-diffusion
  - reasoning
  - math
  - code

CEDR: Continuous Embedding Diffusion for Reasoning

CEDR generates reasoning solutions with continuous latent diffusion and a learned prompt encoder. This repository contains six selected-weight inference packages for GSM8K, MATH500, and coding, before and after NFT.

CEDR uses Qwen3-4B-Instruct-2507 as a frozen representation teacher during training and reuses its tokenizer and frozen token embeddings. CEDR’s diffusion backbone and prompt encoder are trained separately.

Paper: Reasoning with Continuous Latent Diffusion.

The training and inference code and installation instructions are at chengxiang/CEDR.

Package Selected model / prompt EMA Steps SCCFG
gsm8k-b-pre-nft .9999 / .9999 64 2
gsm8k-b-post-nft .99 / .99 64 2
math-l-pre-nft .9999 / .9999 64 3
math-l-post-nft .99 / .99 64 3
oci-l-pre-nft .9999 / .9999 128 3
oci-l-post-nft .9 / .9 128 3

All use async clocks (2.5, 2, 1.5), CFG 1, batch size 32, and a 1,024-token diffusion canvas. The GSM8K post-NFT preset caps decoded responses at 767 tokens. Other presets use the canvas available after each prompt.

Usage

pip install 'git+https://github.com/chengxiang/CEDR.git#egg=cedr[data,math]'
cedr download --model gsm8k-b-pre-nft --output artifacts/gsm8k-b-pre
cedr prepare-benchmark --benchmark gsm8k --checkpoint artifacts/gsm8k-b-pre --output artifacts/gsm8k-test
cedr generate --checkpoint artifacts/gsm8k-b-pre --data artifacts/gsm8k-test --seed 42 --output artifacts/seed42
cedr score --config artifacts/gsm8k-b-pre/config.json --data artifacts/gsm8k-test --predictions artifacts/seed42/predictions.jsonl --output artifacts/seed42/scoring

See the code repository's evaluation guide for MATH500 and official HumanEval/MBPP base/plus scoring. Coding generation uses the combined 164 HumanEval and 378 MBPP tasks. Standard scoring does not apply function-name repair. The optional --function-name-repair flag produces the separately reported alias-scored comparison; see the evaluation guide.

Contents and numerical settings

Each package has FP32 learned denoiser/vocabulary-decoder weights and prompt-encoder weights in safetensors, architecture and inference configuration, selected-state metadata, and SHA-256 hashes. Keep the matched components together. The single shared BF16 token embedding table and tokenizer come from Qwen3-4B-Instruct-2507, revision cdbee75f17c01a7cc42f958dc650907174af0554. Generation requires no teacher Transformer or training cache. Each package also references a matched representation.pt for joint fine-tuning.

The initial noise stream is CPU-seeded. Euler sampling uses a logit-normal quantile grid with mean -1.5 and standard deviation .8, local clock derivatives, and recurrent self-conditioning. Learned weights and ODE states remain FP32; denoiser Transformer forwards use BF16 and vocabulary projection uses FP32. Math/GSM prompt encoders run in FP32; coding prompt encoders run in BF16 with lengths padded to a multiple of 32.

Generation supports resumption after completed batches. Downloads resolve an immutable Hub revision and verify every file. Use --revision to select a recorded release commit.

Joint fine-tuning

The three bundles under representations/gsm8k-b, representations/math-l, and representations/oci-l contain the exact FP32 centering means, effective whitened encoders, and reconstruction matrices. Each family shares its bundle across pre-NFT and post-NFT models. cedr download places the matched bundle in the model directory and verifies its hash.

Use cedr train joint --init-weights MODEL_DIR to initialize the downloaded ELF/prompt pair with fresh optimizers and independent EMAs. Download the pinned full Qwen teacher separately to encode your training data, using either live extraction or an optional feature cache. Separate covariance artifacts are unnecessary when keeping the representation fixed. See the fine-tuning guide.

License and attribution

Model artifacts are Apache 2.0. The CEDR code is MIT. The frozen embedding table and tokenizer retain Qwen attribution; see LICENSE and NOTICE.

These are task-specific research models. Benchmark correctness is measured by the numerical, symbolic, or execution-based scorer specified for each task.

Release validation

Seed-42 checks; separate from the paper's eight-seed estimates. All predictions and per-problem correctness decisions match.

Model Benchmark Correct / total Pass@1 Plus correct / total Plus pass@1
gsm8k-b-pre-nft gsm8k 523/1319 39.65% — —
gsm8k-b-post-nft gsm8k 554/1319 42.00% — —
math-l-pre-nft math500 103/500 20.60% — —
math-l-post-nft math500 128/500 25.60% — —
oci-l-pre-nft humaneval 49/164 29.88% 46/164 28.05%
oci-l-pre-nft mbpp 70/378 18.52% 61/378 16.14%
oci-l-post-nft humaneval 54/164 32.93% 53/164 32.32%
oci-l-post-nft mbpp 92/378 24.34% 80/378 21.16%