shibatch's picture
Upload folder using huggingface_hub
7cdce5b verified
|
Raw
History Blame Contribute Delete
8.15 kB
metadata
library_name: pytorch
pipeline_tag: text-to-image
license: mit
datasets:
  - ylecun/mnist
tags:
  - diffusion
  - ddim
  - mnist
  - handwritten-digits
  - tiny-model

Tiny Digit Diffusion 3M

Tiny Digit Diffusion 3M is a 2,767,529-parameter conditional diffusion model that renders a prompt containing one to eight ASCII digits as a handwritten grayscale image. It was trained from scratch on dynamically composed MNIST digit strings and contains no weights from another generative model.

This is a deliberately narrow tiny-model experiment, not a general-purpose text-to-image model. Its complete prompt language is a numeric string such as 7, 2026, or 31415926.

Generated samples

Repository contents

  • model/model.safetensors: FP32 EMA inference weights
  • model/config.json: architecture and diffusion configuration
  • tiny_digit_diffusion.py: model, schedule, DDIM sampler, and weight loader
  • generate_tiny_digits.py: command-line generation example
  • train_tiny_digit_diffusion.py: training script
  • artifact_metadata.json: training settings and evaluation summary
  • sample_grid.png: generations from the final checkpoint

Optimizer checkpoints, MNIST files, per-epoch samples, and the training log are not included in the distribution package.

Quick start

Create an environment and install the inference dependencies:

python -m venv .venv
source .venv/bin/activate
pip install -r requirements.txt

Generate an eight-digit image:

python generate_tiny_digits.py 31415926 --output 31415926.png

On the first run, the script downloads model/config.json and model/model.safetensors from shibatch/tinydigitdiffusion3m on Hugging Face. Later runs reuse the local Hugging Face cache.

The script accepts the following useful options:

python generate_tiny_digits.py 2026 \
  --steps 50 \
  --guidance-scale 1.0 \
  --seed 0 \
  --device auto \
  --output 2026.png

Use a different Hub revision with --revision, or override the repository with --repo-id. To use model files already stored locally (including the model/ directory included in this distribution package), pass --model-dir:

python generate_tiny_digits.py 2026 --model-dir ./model --output 2026.png

The output is a fixed 32 x 256 grayscale PNG. Short prompts are centered on the eight available 32-pixel slots.

Python example

import sys
from pathlib import Path

import torch
from huggingface_hub import snapshot_download
from PIL import Image

repo_dir = Path(snapshot_download(
    repo_id="shibatch/tinydigitdiffusion3m",
    allow_patterns=[
        "model/config.json",
        "model/model.safetensors",
        "tiny_digit_diffusion.py",
    ],
))
sys.path.insert(0, str(repo_dir))

from tiny_digit_diffusion import ddim_sample, load_model

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = load_model(repo_dir / "model", device)

image = ddim_sample(
    model,
    prompts=["31415926"],
    device=device,
    sampling_steps=50,
    guidance_scale=1.0,
    seed=0,
)[0, 0]

pixels = ((image.cpu() + 1) * 127.5).round().clamp(0, 255).byte().numpy()
Image.fromarray(pixels, mode="L").save("31415926.png")

Prompts must contain only one to eight ASCII digits. Spaces, signs, decimal points, letters, and strings longer than eight digits are rejected.

Recommended inference settings

sampling_method: deterministic DDIM
sampling_steps: 50
guidance_scale: 1.0
output_size: 32 x 256

Use guidance_scale=1.0. Larger classifier-free-guidance values are not helpful for this checkpoint. In particular, the original diagnostic setting of 3.0 over-amplifies the conditioning difference and creates digit-like artifacts in padding slots. At 1.0, the sampler uses the conditional model directly and skips the unnecessary unconditional pass.

Architecture

model: custom conditional diffusion U-Net
parameter_count: 2,767,529
weight_dtype: float32
weight_file_size: approximately 11 MB

image_channels: 1
image_height: 32
image_width: 256
maximum_digits: 8
slot_width: 32

base_channels: 48
channel_multipliers: [1, 1.5, 2, 2.6667]
embedding_dim: 192
token_embedding_dim: 32
spatial_condition_channels: 8
attention_heads: 4
diffusion_steps: 400

The denoiser has four resolution levels with residual blocks, skip connections, a bottleneck self-attention layer, sinusoidal timestep conditioning, global prompt conditioning, and slot-aligned spatial digit conditioning. Prompt tokens consist of digits 0 through 9, a padding token, and a classifier-free null token.

Training

Each training example was assembled dynamically from MNIST:

  1. Sample a prompt length uniformly from one to eight digits.
  2. Sample an MNIST image independently for every digit.
  3. Center the string on a 32 x 256 canvas.
  4. Apply small horizontal and vertical offsets and an intensity variation.
  5. Train the model to predict noise at a random cosine-schedule timestep.

Training configuration:

dtype: float32
epochs: 30
steps: 28,110
batch_size: 64
samples_processed: 1,799,040
optimizer: AdamW
learning_rate: 2.0e-4
warmup_steps: 500
minimum_learning_rate: 2.0e-5
weight_decay: 0.0
gradient_clip: 1.0
condition_dropout: 0.1
ema_decay: 0.999
final_recent_noise_mse: 0.015345
training_time: approximately 70 minutes

The distributed checkpoint is the EMA model, not the raw final optimizer weights.

To train from scratch, install the additional dependencies and choose an output directory:

pip install -r requirements-train.txt
python train_tiny_digit_diffusion.py \
  --output-dir runs/tiny_digit_diffusion_3m \
  --data-dir data/mnist \
  --epochs 30 \
  --batch-size 64 \
  --device cuda

Evaluation

A separate small MNIST classifier was trained only for automated readability measurement. It reached 98.68% accuracy on the MNIST test split. The final generator was then evaluated on 240 random prompts: 30 prompts for every length from one through eight digits, using 50 DDIM steps and guidance 1.0.

Prompt length Exact string Per-digit accuracy
1 29/30 (96.7%) 96.7%
2 26/30 (86.7%) 91.7%
3 26/30 (86.7%) 95.6%
4 22/30 (73.3%) 93.3%
5 22/30 (73.3%) 93.3%
6 14/30 (46.7%) 88.9%
7 22/30 (73.3%) 95.7%
8 13/30 (43.3%) 90.4%
Overall 174/240 (72.5%) 92.6%

These are OCR-proxy measurements, not human ratings. A generated digit can be legible to a person while being classified differently, and the classifier can also be confidently wrong. The samples should therefore be considered alongside the numeric results.

Limitations

  • Longer strings compound individual digit errors; eight-digit exact match is substantially lower than single-digit accuracy.
  • Some handwritten 3, 5, 7, 8, and 9 shapes can be ambiguous.
  • The output canvas and maximum string length are fixed by the architecture.
  • The model supports only ASCII digits and does not render signs, punctuation, decimal values, mathematical expressions, or arbitrary text.
  • This checkpoint reproduces MNIST-like handwriting only. It is not suitable for OCR security testing, document generation, or realistic typography.
  • Results vary with the random seed even though each DDIM trajectory is deterministic for a fixed seed.

Scope

This checkpoint is intended for education, architecture experiments, tests, and demonstrations of a complete conditional image generator at a very small parameter count. It should not be interpreted as an official MNIST model or as an image-generation counterpart of any production diffusion system.

License

The source code and model checkpoint in this package are released under the MIT License. See LICENSE for the complete terms.

The MNIST images are not redistributed in this package. Training used MNIST through torchvision; the MNIST dataset card identifies its license as MIT and credits Yann LeCun, Corinna Cortes, and Christopher J. C. Burges as the dataset curators.