tchauffi commited on
Commit
0f956c4
·
verified ·
1 Parent(s): e337120

Update SudokuDiT card + assets

Browse files
Files changed (5) hide show
  1. .gitattributes +1 -0
  2. README.md +102 -0
  3. config.json +9 -0
  4. model.safetensors +3 -0
  5. solve.gif +3 -0
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ solve.gif filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,102 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ pipeline_tag: other
4
+ library_name: pytorch
5
+ tags:
6
+ - maze
7
+ - path-planning
8
+ - masked-diffusion
9
+ - diffusion-transformer
10
+ - dit
11
+ - reasoning
12
+ datasets:
13
+ - sapientinc/maze-30x30-hard-1k
14
+ ---
15
+
16
+ # MazeDiT-30x30
17
+
18
+ A **1.38 M-parameter** Diffusion Transformer that solves 30x30 maze path-planning as
19
+ **masked discrete diffusion**. It labels *every* open cell as on-path or off-path, most
20
+ confident first — it never traces a route.
21
+
22
+ ![watch it solve](solve.gif)
23
+
24
+ *A 131-cell shortest path recovered in 40 adaptive steps. Walls are near-black, cells not
25
+ yet labelled stay slate, and path cells are tinted by the model's step-0 confidence
26
+ (blue = unsure -> green = sure).*
27
+
28
+ Same recipe as [`tchauffi/sudoku-dit`](https://huggingface.co/tchauffi/sudoku-dit), with the
29
+ 9-digit vocabulary swapped for 3 tokens (wall / open / path).
30
+
31
+ - **Architecture:** GridDiT — DiT with adaLN-Zero conditioning, `hidden=128`,
32
+ `heads=4`, `blocks=4`; per-cell token + 2-D positional embeddings, plus a
33
+ timestep. No Sudoku box embedding.
34
+ - **Code, training and an interactive web demo:** <https://github.com/tchauffi/nonet>
35
+
36
+ ## Input / output
37
+
38
+ A maze is **900 tokens, row-major**: `0` = `[MASK]` (an open cell to label), `1` = wall,
39
+ `2` = open/off-path, `3` = path. The question marks walls and the two endpoints (start and
40
+ goal, both token `3`) and masks everything else; the solver clamps the givens.
41
+
42
+ ## Usage
43
+
44
+ ```python
45
+ import torch
46
+ from nonet.hub import load_maze_solver # pip install git+https://github.com/tchauffi/nonet
47
+
48
+ solver = load_maze_solver("tchauffi/maze-dit") # config records the cosine-high schedule
49
+ question = torch.tensor([[...]]) # (1, 900), see the encoding above
50
+ pred = solver.solve(question, num_steps=900, conf_threshold=0.999)
51
+ ```
52
+
53
+ `scripts/eval_maze30.py` in the repo downloads the benchmark and reproduces the table below.
54
+
55
+ ## The metric matters here
56
+
57
+ Shortest paths on these mazes are **massively non-unique** — a median of ~10^7 distinct
58
+ optimal routes per maze. Grading against the dataset's single reference answer therefore
59
+ measures *tie-break mimicry*, not solving. We report **`valid_shortest`**: the prediction is
60
+ a valid simple start-to-goal path **and** its length equals the BFS optimum. Both are
61
+ checkable from the question alone, without the reference.
62
+
63
+ ## Performance
64
+
65
+ Full 1,000-maze test split, adaptive decoder (tau = 0.999):
66
+
67
+ | metric | value |
68
+ |--------|-------|
69
+ | valid_shortest (**the honest metric**) | **52.4 %** |
70
+ | valid simple S->G path | 61.8 % |
71
+ | + restart sampling, k = 32 | **57.2 %** valid_shortest |
72
+ | exact match vs reference | 6.9 % |
73
+
74
+ The gap between 61.8 % valid and 6.9 % exact is the degeneracy above: the model routinely
75
+ finds *a* correct path that is not the one the generator happened to emit.
76
+
77
+ For reference, **HRM reports 74.5 %** on this benchmark at **27 M parameters** (~20x larger)
78
+ under the exact-match protocol, which the degeneracy finding makes hard to compare directly.
79
+
80
+ Restart sampling climbs slowly here (52.4 -> 57.2 % over 32 attempts) and is still
81
+ unplateaued — nothing like the near-doubling the same trick gives on Sudoku-Extreme. That
82
+ says the residual failures are a systematic data ceiling, not decoding luck.
83
+
84
+ ## Training
85
+
86
+ - Data: the **1,000 training mazes** provided by
87
+ [`sapientinc/maze-30x30-hard-1k`](https://huggingface.co/datasets/sapientinc/maze-30x30-hard-1k),
88
+ plus dihedral x8 augmentation (the square's symmetry group).
89
+ - Objective: masked cross-entropy over masked cells, conditional (walls and endpoints are
90
+ never masked), **cosine-high** masking schedule.
91
+ - ~18 k steps at the selected checkpoint, AdamW, batch 256, cosine lr decay.
92
+
93
+ ## Limitations
94
+
95
+ - **Checkpoint selection** used a 256-maze subset of the same test split; no separate
96
+ validation split ships with the benchmark (HRM's protocol shares this).
97
+ - **Memorization collapse:** dihedral x8 over 1,000 mazes yields only ~8 k effective boards.
98
+ Past ~20 k steps the model memorizes them — a 60 k-step run drives train cell-accuracy to
99
+ 1.000 and test-valid to 0 %. This checkpoint is the step-18k peak; longer training is
100
+ strictly worse.
101
+ - Trained only on 30x30 mazes of this generator; the released weights regenerate `pos_embed`
102
+ for the grid size, but nothing about other sizes is tested.
config.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "hidden_size": 128,
3
+ "num_heads": 4,
4
+ "mlp_ratio": 4.0,
5
+ "num_blocks": 4,
6
+ "grid_size": 30,
7
+ "num_classes": 3,
8
+ "schedule": "cosine-high"
9
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:abe29b8c9ddc250b150d995b724491d813e4ab8d5676d4d115769f0cb22d0a2c
3
+ size 5541548
solve.gif ADDED

Git LFS Details

  • SHA256: 3a1e03194ed3a449236bc7ccbddcaf7b865db9babd2375798ee030582fd2abac
  • Pointer size: 131 Bytes
  • Size of remote file: 554 kB