maze-dit / README.md
tchauffi's picture
Restore MazeDiT weights (previous push uploaded a Sudoku model by mistake)
2716eba verified
|
Raw
History Blame Contribute Delete
4.29 kB
---
license: apache-2.0
pipeline_tag: other
library_name: pytorch
tags:
- maze
- path-planning
- masked-diffusion
- diffusion-transformer
- dit
- reasoning
datasets:
- sapientinc/maze-30x30-hard-1k
---
# MazeDiT-30x30
A **1.38 M-parameter** Diffusion Transformer that solves 30x30 maze path-planning as
**masked discrete diffusion**. It labels *every* open cell as on-path or off-path, most
confident first β€” it never traces a route.
![watch it solve](solve.gif)
*A 131-cell shortest path recovered in 40 adaptive steps. Walls are near-black, cells not
yet labelled stay slate, and path cells are tinted by the model's step-0 confidence
(blue = unsure -> green = sure).*
Same recipe as [`tchauffi/sudoku-dit`](https://huggingface.co/tchauffi/sudoku-dit), with the
9-digit vocabulary swapped for 3 tokens (wall / open / path).
- **Architecture:** GridDiT β€” DiT with adaLN-Zero conditioning, `hidden=128`,
`heads=4`, `blocks=4`; per-cell token + 2-D positional embeddings, plus a
timestep. No Sudoku box embedding.
- **Code, training and an interactive web demo:** <https://github.com/tchauffi/nonet>
## Input / output
A maze is **900 tokens, row-major**: `0` = `[MASK]` (an open cell to label), `1` = wall,
`2` = open/off-path, `3` = path. The question marks walls and the two endpoints (start and
goal, both token `3`) and masks everything else; the solver clamps the givens.
## Usage
```python
import torch
from nonet.hub import load_maze_solver # pip install git+https://github.com/tchauffi/nonet
solver = load_maze_solver("tchauffi/maze-dit") # config records the cosine-high schedule
question = torch.tensor([[...]]) # (1, 900), see the encoding above
pred = solver.solve(question, num_steps=900, conf_threshold=0.999)
```
`scripts/eval_maze30.py` in the repo downloads the benchmark and reproduces the table below.
## The metric matters here
Shortest paths on these mazes are **massively non-unique** β€” a median of ~10^7 distinct
optimal routes per maze. Grading against the dataset's single reference answer therefore
measures *tie-break mimicry*, not solving. We report **`valid_shortest`**: the prediction is
a valid simple start-to-goal path **and** its length equals the BFS optimum. Both are
checkable from the question alone, without the reference.
## Performance
Full 1,000-maze test split, adaptive decoder (tau = 0.999):
| metric | value |
|--------|-------|
| valid_shortest (**the honest metric**) | **52.4 %** |
| valid simple S->G path | 61.8 % |
| + restart sampling, k = 32 | **57.2 %** valid_shortest |
| exact match vs reference | 6.9 % |
The gap between 61.8 % valid and 6.9 % exact is the degeneracy above: the model routinely
finds *a* correct path that is not the one the generator happened to emit.
For reference, **HRM reports 74.5 %** on this benchmark at **27 M parameters** (~20x larger)
under the exact-match protocol, which the degeneracy finding makes hard to compare directly.
Restart sampling climbs slowly here (52.4 -> 57.2 % over 32 attempts) and is still
unplateaued β€” nothing like the near-doubling the same trick gives on Sudoku-Extreme. That
says the residual failures are a systematic data ceiling, not decoding luck.
## Training
- Data: the **1,000 training mazes** provided by
[`sapientinc/maze-30x30-hard-1k`](https://huggingface.co/datasets/sapientinc/maze-30x30-hard-1k),
plus dihedral x8 augmentation (the square's symmetry group).
- Objective: masked cross-entropy over masked cells, conditional (walls and endpoints are
never masked), **cosine-high** masking schedule.
- ~18 k steps at the selected checkpoint, AdamW, batch 256, cosine lr decay.
## Limitations
- **Checkpoint selection** used a 256-maze subset of the same test split; no separate
validation split ships with the benchmark (HRM's protocol shares this).
- **Memorization collapse:** dihedral x8 over 1,000 mazes yields only ~8 k effective boards.
Past ~20 k steps the model memorizes them β€” a 60 k-step run drives train cell-accuracy to
1.000 and test-valid to 0 %. This checkpoint is the step-18k peak; longer training is
strictly worse.
- Trained only on 30x30 mazes of this generator; the released weights regenerate `pos_embed`
for the grid size, but nothing about other sizes is tested.