| --- |
| license: mit |
| tags: |
| - diffusion |
| - ddpm |
| - text-to-image |
| - pytorch |
| - fashion-mnist |
| - classifier-free-guidance |
| library_name: pytorch |
| pipeline_tag: text-to-image |
| --- |
| |
| # text-diffusion-fashion-mnist |
|
|
| A **text-conditioned diffusion model built from scratch in PyTorch** — a miniature Stable |
| Diffusion trained on Fashion-MNIST (32×32 grayscale). A U-Net learns to reverse a Gaussian |
| noising process, conditioned on **frozen CLIP text embeddings**, and generates a garment image |
| from a caption using **classifier-free guidance**. |
|
|
| 📦 **Code / training / full write-up:** https://github.com/mbsdeepak/text-diffusion-fashion-mnist |
|
|
|  |
|
|
| *One row per class (t-shirt, trouser, pullover, dress, coat, sandal, shirt, sneaker, bag, ankle |
| boot); each image is generated from pure noise, DDIM 50 steps, guidance 1.5.* |
|
|
| ## What it does |
|
|
| Give it one of the 10 Fashion-MNIST categories and it synthesises a brand-new image of that |
| item from random noise: |
|
|
| ``` |
| "sneaker" ──► [model] ──► a novel 32×32 image of a sneaker |
| ``` |
|
|
| ## Files |
|
|
| | File | Description | |
| |------|-------------| |
| | `model.safetensors` | U-Net weights (19.4M params, raw / non-EMA) | |
| | `config.json` | The `Config` used to build the U-Net | |
|
|
| ## How to load & sample |
|
|
| The architecture is defined in the [GitHub repo](https://github.com/mbsdeepak/text-diffusion-fashion-mnist), |
| so load the weights into it: |
|
|
| ```bash |
| git clone https://github.com/mbsdeepak/text-diffusion-fashion-mnist |
| cd text-diffusion-fashion-mnist |
| pip install -r requirements.txt |
| huggingface-cli download mbsdeepak/text-diffusion-fashion-mnist model.safetensors --local-dir . |
| ``` |
|
|
| ```python |
| import torch |
| from safetensors.torch import load_file |
| from config import get_config, FASHION_CLASSES |
| from src.unet import UNet |
| from src.diffusion import GaussianDiffusion |
| from src.text_encoder import TextConditioner |
| from src.data import denormalize |
| from torchvision.utils import save_image |
| |
| cfg = get_config() |
| model = UNet(cfg).to(cfg.device) |
| model.load_state_dict(load_file("model.safetensors")) |
| model.eval() |
| |
| cond = TextConditioner(cfg).to(cfg.device) |
| diff = GaussianDiffusion(cfg).to(cfg.device) |
| |
| labels = torch.arange(len(FASHION_CLASSES), device=cfg.device) # one of each class |
| imgs = diff.ddim_sample(model, cond, labels) |
| save_image(denormalize(imgs), "out.png", nrow=len(FASHION_CLASSES)) |
| ``` |
|
|
| ## Training details |
|
|
| - **Data:** Fashion-MNIST, 32×32, normalized to [-1, 1] |
| - **Objective:** ε-prediction MSE (DDPM), cosine noise schedule, T=1000 |
| - **Conditioning:** frozen CLIP (`openai/clip-vit-base-patch32`) text embeddings via FiLM + |
| cross-attention; 15% caption dropout for classifier-free guidance |
| - **Trained:** 25 epochs on Apple Silicon (MPS), final loss ≈ 0.042 |
| - **Sampling:** DDIM, 50 steps, guidance scale 1.5 (low guidance — high CFG artifacts on this |
| small, short-trained model) |
|
|
| > These are the **raw** weights, not EMA — for a short (25-epoch) run the EMA average still lags |
| > the live weights, so the raw model produces the cleaner samples. |
|
|
| ## Limitations |
|
|
| - Conditioned on a **fixed set of 10 captions**, so it's text-driven class-conditional |
| generation, not open-vocabulary text-to-image (the cross-attention plumbing is the same; the |
| dataset is the limit). |
| - 32×32 grayscale — a learning/portfolio project to understand the method, not photorealism. |
|
|
| ## References |
|
|
| Ho et al. *DDPM* (2020) · Nichol & Dhariwal *Improved DDPM* (2021) · Song et al. *DDIM* (2021) · |
| Ho & Salimans *Classifier-Free Guidance* (2022) · Rombach et al. *Latent Diffusion* (2022). |
|
|