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
---