File size: 5,580 Bytes
edf255a 4649ea9 edf255a 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 352eec2 4649ea9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | ---
license: apache-2.0
pipeline_tag: text-to-image
---
# ShellD (Shell Diffusion)
**Small DiT-based Text-to-Image Latent Diffusion Model**
ShellD is a lightweight text-to-image model that generates **256×256** images from natural language prompts. It uses a Diffusion Transformer (DiT) backbone operating in a compact VAE latent space, making it feasible to train and run on consumer GPUs.
---
## Model Architecture
```
Text Prompt → [MiniLM-L6-v2 (frozen)] → Text Embedding (384-d)
↓
Random Noise → [VAE Encoder] → Latent (16ch) → [DiT (12 blocks)] → Denoised Latent → [VAE Decoder] → 256×256 Image
```
| Component | Details | Params |
|-----------|---------|--------|
| **Text Encoder** | `sentence-transformers/all-MiniLM-L6-v2` (frozen) | 22.71M |
| **VAE** | Encoder + Decoder with residual blocks, 3 down/up stages, latent dim=16 | 23.43M |
| **DiT** | 12-layer Transformer with self-attention, cross-attention (text), and adaptive timestep conditioning. Patch size=4, hidden dim=256, 8 heads | 20.80M |
| **Total** | | **66.95M** (trainable: **44.23M**) |
### VAE (Autoencoder)
The VAE compresses 256×256 RGB images into a **16-channel latent** with spatial size 32×32 (downsampled by 8×). It uses residual blocks with GroupNorm and SiLU activations. During training, a KL penalty (β=0.1) keeps latents close to a standard normal distribution.
### DiT (Diffusion Transformer)
The DiT operates on patched latents (patch size 4 → 8×8 = 64 patches). Each block includes:
- **Self-attention** for spatial relationships
- **Cross-attention** conditioned on text embeddings
- **Adaptive timestep conditioning** via an MLP-projected sinusoidal embedding
- **Dropout** (0.1) in attention and MLP for regularization
### Diffusion Process
Standard DDPM (Denoising Diffusion Probabilistic Model) with 1000 timesteps and a linear beta schedule (β₁=1e‑4, βᵀ=0.02). The model is trained to predict the added noise ε. Classifier-free guidance (CFG) is used during training with a text-conditioning dropout probability of 15%.
---
## Training
### Dataset
**[jackyhate/text-to-image-2M](https://huggingface.co/datasets/jackyhate/text-to-image-2M)** — \~2M high-quality text-image pairs in webdataset format. Loaded via the `datasets` library with **streaming** to avoid materializing the full 2TB+ dataset into memory. A rotating **2000-image in-memory buffer** (\~400 MB RAM) is refreshed each epoch from a fresh random stream to provide shuffle diversity without disk I/O bottlenecks.
### Data Augmentation
Random horizontal flip (p=0.5), color jitter (brightness/contrast/saturation ±0.2, hue ±0.05), and random affine transforms (rotation ±10°, translation ±5%, scale 0.9–1.1) via `torchvision`.
### Hyperparameters
| Parameter | Value |
|-----------|-------|
| **Image size** | 256×256 |
| **Batch size** | 8 |
| **Optimizer** | AdamW (β₁=0.9, β₂=0.999, lr=1e‑4, weight decay=0.01) |
| **LR schedule** | Linear warmup (500 steps) + Cosine annealing |
| **Gradient clipping** | 1.0 (norm) |
| **Mixed precision** | FP16 via `torch.cuda.amp.GradScaler` |
| **Dropout** | 0.1 (DiT attention + MLP) |
| **EMA** | Exponential moving average (decay=0.999) applied at every step; EMA weights used for validation and final checkpoint |
| **Early stopping** | Patience of 8 epochs on validation loss (10% held-out split) |
| **KL weight** | 0.1 (β-VAE style) |
### VAE Pretraining
Before diffusion training, the VAE is pretrained for **10 epochs** on reconstruction + KL loss with a higher learning rate (lr=1e‑3) to establish a meaningful latent space. During this phase the DiT and text encoder are frozen.
### Training Phases
1. **VAE Pretraining** (10 epochs) — Train encoder + decoder on image reconstruction to establish a meaningful latent space. DiT and text encoder are frozen.
2. **DiT Diffusion Training** (up to 30 epochs, early-stopped) — Freeze VAE, train DiT to denoise latents conditioned on text embeddings. CFG dropout randomly replaces text embeddings with a learned null embedding to enable classifier-free guidance at inference time.
---
## Usage
### Requirements
```bash
pip install torch safetensors sentence-transformers pillow numpy huggingface-hub
```
For training, also install:
```bash
pip install torchvision datasets
```
### Inference (standalone — loads from Hugging Face)
```python
from inference import ShellDInference
pipe = ShellDInference("FlameF0X/ShellD")
image = pipe.generate("a serene lake surrounded by mountains")
image.save("output.png")
```
`ShellDInference` automatically downloads weights from Hugging Face via `huggingface_hub` on first use and caches them locally.
### Streaming Generation
View the diffusion process unfold step-by-step:
```python
for img, step_info in pipe.generate_stream(
prompt="a futuristic city at night",
num_steps=250,
cfg_scale=3.0,
display_every=25, # emit an image every 25 steps
):
print(f"Step {step_info['step']}/{step_info['total']}")
img.save(f"progress_{step_info['step']:04d}.png")
```
---
## Intended Use
- Educational exploration of diffusion transformers
- Lightweight text-to-image generation on consumer hardware
- Starting point for fine-tuning on custom datasets
## Limitations
- 256×256 resolution only (no upscaling built in)
- Limited prompt understanding due to small DiT and frozen lightweight text encoder
- Quality depends on training data distribution — may not match large-scale models like SDXL or Flux
--- |