Team Ai
Modelpublic

FlameF0X/ShellD

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
5likes30downloads
Model Card

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
ComponentDetailsParams
Text Encodersentence-transformers/all-MiniLM-L6-v2 (frozen)22.71M
VAEEncoder + Decoder with residual blocks, 3 down/up stages, latent dim=1623.43M
DiT12-layer Transformer with self-attention, cross-attention (text), and adaptive timestep conditioning. Patch size=4, hidden dim=256, 8 heads20.80M
Total66.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

ParameterValue
Image size256×256
Batch size8
OptimizerAdamW (β₁=0.9, β₂=0.999, lr=1e‑4, weight decay=0.01)
LR scheduleLinear warmup (500 steps) + Cosine annealing
Gradient clipping1.0 (norm)
Mixed precisionFP16 via torch.cuda.amp.GradScaler
Dropout0.1 (DiT attention + MLP)
EMAExponential moving average (decay=0.999) applied at every step; EMA weights used for validation and final checkpoint
Early stoppingPatience of 8 epochs on validation loss (10% held-out split)
KL weight0.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. 1.VAE Pretraining (10 epochs) — Train encoder + decoder on image reconstruction to establish a meaningful latent space. DiT and text encoder are frozen.
  2. 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 ---