This project implements a simplified RandAR-style random-order autoregressive Transformer on the MNIST dataset.
The goal is to understand the first principles behind random-order autoregressive generation and test whether a model trained only on randomly ordered complete images can generalize to tasks such as partial-image completion and inpainting and outpainting.
For an image represented by pixels
the joint probability can be factorized using the probability chain rule:
where
A conventional autoregressive image model normally uses one fixed order, such as raster order:
RandAR instead samples a different random permutation during training:
The joint distribution is unchanged, but the model learns many different conditional distributions.
Because pixel order is random, the model must know which location it is currently being asked to predict.
For each pixel
A training sequence therefore looks like
Conceptually,
means:
What is the pixel value at location (i)?
The model predicts
At each position instruction, the Transformer predicts a probability distribution over possible pixel values.
The RandAR loss is negative log-likelihood:
In the MNIST implementation, grayscale values are discretized into intensity bins. Therefore the loss is implemented using categorical cross-entropy:
for the discrete pixel-value distribution.
The complete training objective can be viewed as
Random ordering exposes the model to many different conditioning sets.
For example, the same pixel
or
or other combinations depending on the sampled permutation.
Therefore the model learns the more general task
where (S) can be many different subsets of known pixels.
This is the main source of RandAR's flexibility.
At inference, any known subset of pixels can be placed first in the autoregressive context.
Suppose
The model generates
autoregressively:
Each newly sampled pixel becomes additional context for the next prediction.
This allows the same model to perform:
- random missing-pixel completion,
- center-region inpainting,
- left/right-half completion,
- structured-mask completion,
- stochastic generation of multiple plausible reconstructions.
Importantly, the model is trained only with random-order complete images. Structured masks are introduced only during inference.
A standard autoregressive Transformer normally assumes a fixed sequence order:
RandAR separates two concepts:
from
The generation order is determined by the random permutation
Thus RandAR is still fully autoregressive and causal, but it is not restricted to one fixed autoregressive factorization.
The MNIST implementation uses a small decoder-only causal Transformer.
Main components:
- MNIST
$(28\times28)$ grayscale images; - 784 pixel locations;
- discretized grayscale pixel-value tokens;
- row and column embeddings for spatial location;
- explicit position-instruction embeddings;
- pixel-value embeddings;
- causal self-attention;
- Transformer blocks with LayerNorm and MLP layers;
- categorical output head over pixel-intensity bins;
- KV caching for faster autoregressive sampling.
The model sequence is approximately
under a different random ordering for every training example.
The entire model can be summarized as
The important consequence is that one model learns many conditional views of the same joint image distribution rather than being tied to one fixed generation order.