ReMDM Planner: Craftax checkpoints
Trained weights accompanying Return-Weighted ELBO Fine-Tuning Degrades Masked Diffusion Planners: a remasking discrete diffusion model (ReMDM) used as an action-sequence planner in Craftax, together with the PPO-RNN experts that supervise it, and the results reported in the paper.
Code, configs and evaluation harness: https://github.com/mathisweil/craftax-ReMDM-planner
Contents
| Path | Role | Environment | Architecture | Selected at | Training | Size |
|---|---|---|---|---|---|---|
checkpoints/offline/Craftax-Classic-Symbolic-v1-Offline-Diffusion-BC-100M |
Diffusion planner (offline BC) | Craftax-Classic-Symbolic-v1 |
6L, d_model 384, 8 heads, horizon 32 | 99,942,400 | 97,600 grad steps | 97 MB |
checkpoints/online/Craftax-Classic-Symbolic-v1-Online-Diffusion-DAgger-100M |
Diffusion planner (online DAgger) | Craftax-Classic-Symbolic-v1 |
6L, d_model 384, 8 heads, horizon 32 | 40,370,176 | 97,600 grad steps | 33 MB |
checkpoints/ppo_agents/Craftax-Classic-Symbolic-v1-PPO_RNN-1000M |
PPO-RNN expert | Craftax-Classic-Symbolic-v1 |
RNN, layer size 512 | 1,000,000,000 | 1e+09 frames | 35 MB |
checkpoints/ppo_agents/Craftax-Symbolic-v1-PPO_RNN-1000M |
PPO-RNN expert | Craftax-Symbolic-v1 |
RNN, layer size 512 | 1,000,000,000 | 1e+09 frames | 50 MB |
Each diffusion checkpoint ships a resume_metadata.json holding the full
config snapshot it was trained under; each PPO expert ships config.yaml and
wandb-summary.json (final training metrics).
Weights are Orbax checkpoint directories
(OCDBT format), not safetensors — the models are Flax modules restored via
orbax.checkpoint, and the paths above mirror the source repository so a
snapshot can be dropped straight into a working copy.
Results
RL fine-tuning ablation runs, as produced by experiments/rl_finetuning/run_ablations.py. Each run ships its results.json summary, the diagnosis.md write-up, and the tables (.csv and .tex) and figures generated from it.
| Run | Contents | Size |
|---|---|---|
experiments/rl_finetuning/outputs/craftax_classic_ablations |
results.json, diagnosis.md, 22 tables, 113 figures, 4 gdelta |
37 MB |
experiments/rl_finetuning/outputs/review_anchor_baseline_rl |
results.json, diagnosis.md, 18 tables, 16 figures |
2 MB |
experiments/rl_finetuning/outputs/review_run1_bc_all |
results.json, diagnosis.md, 18 tables, 16 figures |
2 MB |
experiments/rl_finetuning/outputs/review_run2_advclip_lr_matched |
results.json, diagnosis.md, 18 tables, 16 figures |
2 MB |
experiments/rl_finetuning/outputs/review_run4_baseline_lr1e-4 |
results.json, diagnosis.md, 18 tables, 16 figures |
2 MB |
experiments/rl_finetuning/outputs/review_run4_baseline_lr1e-5 |
results.json, diagnosis.md, 18 tables, 16 figures |
2 MB |
Evaluation results produced by main.py --mode inference on the checkpoints above, under results/inference/.
| File | Environment | Evaluation | Headline metric | Size |
|---|---|---|---|---|
eval_classic_bc_s42.json |
Craftax-Classic-Symbolic-v1 |
32 envs x 10000 steps | mean score 3.88 | 1 KB |
eval_classic_dagger_s42.json |
Craftax-Classic-Symbolic-v1 |
32 envs x 10000 steps | mean score 3.26 | 1 KB |
expert_classic_n256_s0_t1.0.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 18.78 | 1 KB |
expert_classic_n32_s0_t0.5.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 18.48 | 1 KB |
expert_classic_n32_s0_t1.0.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 18.38 | 1 KB |
expert_classic_n32_s1_t0.5.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 19.04 | 1 KB |
expert_classic_n32_s1_t1.0.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 18.82 | 1 KB |
expert_classic_n32_s2_t0.5.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 19.10 | 1 KB |
expert_classic_n32_s2_t1.0.json |
Craftax-Classic-Symbolic-v1 |
- | mean score 19.04 | 1 KB |
Manuscript figures, as vector PDF at NeurIPS column width, under results/paper_figures/. These are built by scripts/paper_figures.py, which reads the ablation results.json of both environments and draws Craftax Classic and MiniHack side by side, so the identical set is published in this release and in the MiniHack one.
| Figure | Size |
|---|---|
fig10_timestep_conditioning.pdf |
31 KB |
fig11_train_vs_eval.pdf |
34 KB |
fig1_finetuning_trajectories.pdf |
21 KB |
fig2_repr_drift.pdf |
24 KB |
fig3_cka.pdf |
19 KB |
fig4_grad_alignment.pdf |
23 KB |
fig5_minihack_per_env.pdf |
18 KB |
fig6_achievements.pdf |
24 KB |
fig7_tbin_gradients.pdf |
20 KB |
fig8_score_vs_kl.pdf |
23 KB |
fig9_weight_dispersion.pdf |
45 KB |
Download
This repo mirrors the code repository's layout, so a snapshot drops straight
into a working copy -- but it also carries its own README.md (this card),
LICENSE and .gitattributes, and local_dir="." would overwrite the code
repository's copies of all three. Exclude them, or download into a directory
of its own.
from huggingface_hub import snapshot_download
# everything (~264 MB), into a clone of the code repository
snapshot_download(
repo_id="mathisweil/remdm-craftax-checkpoints",
local_dir=".",
ignore_patterns=["README.md", "LICENSE", ".gitattributes"],
)
# or somewhere of its own, leaving any working copy untouched
snapshot_download(repo_id="mathisweil/remdm-craftax-checkpoints", local_dir="remdm-craftax")
# a single model
snapshot_download(
repo_id="mathisweil/remdm-craftax-checkpoints",
local_dir=".",
allow_patterns="checkpoints/offline/Craftax-Classic-Symbolic-v1-Offline-Diffusion-BC-100M/**",
)
Use
From a clone of the code repository, after downloading into it:
uv run python main.py --mode inference \
--checkpoint checkpoints/offline/Craftax-Classic-Symbolic-v1-Offline-Diffusion-BC-100M \
--output results/inference/eval.json
Programmatic loading uses src.planners.model.load_checkpoint for the
diffusion planners and src.planners.ppo.load_ppo_agent for the experts; both
take the checkpoint directory path and restore the latest step. Architecture
arguments should be read from the checkpoint's own resume_metadata.json
rather than hardcoded.
Training
The diffusion planners are bidirectional transformers that denoise a masked action plan conditioned on the symbolic observation, trained either by offline behaviour cloning on PPO rollouts or by online DAgger against the PPO expert. Model size and horizon differ per run (see the table); the PPO-RNN experts are the Craftax baselines. Exact hyperparameters for every run, including the remasking strategy, schedule and sampling settings, are in the per-checkpoint metadata files listed above, which are the authoritative record.
Directory names encode the environment and the total environment timesteps the
run was trained for. Selected at is whatever each run used as its Orbax step
counter, which is environment frames for the runs published here.
Limitations
These are research artefacts tied to specific Craftax versions and symbolic observation encodings; they are not general-purpose agents and will not transfer to other environments or to pixel observations. Evaluation results and their variance are reported in the paper.
Citation
@inproceedings{remdm-craftax-planner,
title = {Return-Weighted ELBO Fine-Tuning Degrades Masked Diffusion Planners},
author = {Weil, Mathis},
year = {2026},
note = {NeurIPS 2026 Workshop: Beyond Next-Token Prediction}
}
License
MIT, see LICENSE.