SDXS — compact single-stream diffusion transformer (text-to-image)

Single-stream DiT (Krea/FLUX-style blocks) with classic depth: 20 unique blocks applied once (6 in + 8 mid + 6 out = 20 passes, 1.996B). The earlier weight-tied looped-middle experiment (6+8×2+6, 28 passes) was dropped after measurements showed the second pass idled and classic depth matched its quality at ~14% less compute (see History).

Model size

  • Weights: 1.996B = 4.0 GB (bf16) — 20 unique blocks: 6 input + 8 middle + 6 output, applied once.
  • Depth: 20 forward passes.
  • Fits 2×RTX 5090 (32GB) at --batch-size 16 WITHOUT --cache-emb (~24 GB/rank, ~3 s/step).

Architecture

Single-stream DiT (RMSNorm + QK-norm, SwiGLU 8/3, AdaNorm bias modulation, 3D axial RoPE, full attention kv==heads):

hidden 2560  (20 heads × 128, FULL attention — kvheads = 20, no GQA)
blocks: 20 unique = 6 input + 8 middle + 6 output, applied once (20 passes)
text:   Qwen3.5-2B (hidden 2048), 3 equally-spaced hidden layers
        (hidden_states[2, 12, 22]) → text-fusion (n=3) → 2560
patch:  2  (image tokens = latent/2² per axis)
VAE:    AsymmetricAutoencoderKL, 32 latent channels, encoder f8 / decoder f16
        (latents_mean/std applied), 2× upscale → generate at 2× train res

History (why the loop is gone)

A looped middle (6 in + 8 mid ×2 + 6 out, 28 passes with 20 weight-tied blocks, per-iteration loop_emb/loop_norms) was trained end-to-end, but diagnostics showed the second pass collapses to a fixed point: it changed the final output by only ~2.5%, the per-pass signal (loop_emb[1]) had ~0.14% effect, and a coarse-to-fine pooled pass-2 variant was zeroed by the optimizer during training. A 1-image overfit test (1000 steps, 2×RTX 5090) confirmed the classic 20-pass model matches or beats the looped one (loss 0.0213 vs 0.0219; PSNR 28.2 vs 27.1 dB) at ~14% less compute. Conclusion: classic depth — no loops.

The classic checkpoint was converted from the looped one (all weights 1:1; loop_emb folded into the mid blocks' modulation bias; mid_norm from loop_norms.0).

Files

transformer_sdxs.py     # SDXSTransformer + config (classic depth, 20 blocks)
pipeline_sdxs.py        # custom pipeline (text → DiT → VAE)
generate.py             # inference
model_index.json, transformer/config.json
vae/  text_encoder/  tokenizer/  scheduler/

# data & training
dataset.py            # images + .txt -> HF dataset (VAE latents + text + size)
make_model.py         # build a fresh random-init SDXS with a given config
train.py              # teacher-free training (Accelerator, flow-matching shift 5,
                      #   resolution sampler, independent word/caption dropout,
                      #   optional --compile / --attn-backend)
train_distill.py      # FLUX.2-klein (quantized SDNQ) → SDXS distillation
train_legacy.py       # superseded pre-loop trainer (kept, not used)
one_sample_train/     # legacy smoke test

# FLUX.2-klein adapter (our text encoder -> FLUX.2-klein teacher space)
qwen_adapter_2b_4b/   # Qwen3.5-2B [2,12,22] → Qwen3-4B [9,18,27] text adapter

Setup / dependencies

pip install torch diffusers transformers accelerate datasets bitsandbytes einops wandb sdnq

Local models / data (paths used by the code):

What Path
SDXS (student) repo root transformer/, vae/, text_encoder/, tokenizer/, scheduler/
Teacher (quantized SDNQ) /workspace/.hf_home/hub/models--Disty0--FLUX.2-klein-4B-SDNQ-4bit-dynamic/snapshots/<hash>/
Qwen adapter qwen_adapter_2b_4b/adapter.pt
Distill datasets /workspace/ds (auto-merged: alchemist + civitai + ..., HF arrow)

Usage

# 1. build a dataset from a folder of images + .txt captions (saves VAE latents)
python dataset.py

# 2. build a random-init model (classic, e.g. 6+8+6 = 20 passes)
python make_model.py --n-in 6 --n-mid 8 --n-out 6

# 3. teacher-free training (canonical; continue from transformer/ checkpoint)
accelerate launch train.py --ds-path /workspace/ds --batch-size 24 \
    --compile --lr 5e-5 --seed 43 --sample-every-steps 500 \
    --save-every-minutes 60 --wandb   # batch 24 + --compile = ~+55% throughput (2026-09-13 bench)

# distillation from FLUX.2-klein through the Qwen adapter
accelerate launch train_distill.py \
    --ds-path /workspace/ds --adapter qwen_adapter_2b_4b/adapter.pt \
    --model-path transformer --batch-size 16 --lr 1e-4 --lambda-gt 0.01 --wandb

# 4. inference
python generate.py --prompt "a majestic deer in a forest"

Training notes

  • Teacher-free train.py: loss = MSE(student velocity, ground-truth velocity), flow-matching, timestep shift = 5.0. Independent per-step dropouts: --caption-dropout 0.05 (whole prompt) + --word-dropout 0.05 (one word). shift 5.0 = the SD3 rule for ~320-640 latents (√(32·40·80/4096) ≈ 5.0); --t-detail-bias 0.3 skews sampling toward low t (detail region).
  • Distillation (train_distill.py): teacher is frozen SDNQ-quantized FLUX.2-klein; loss = MSE(v_student, v_teacher) + lambda_gt·MSE(v_student, v_gt), with --lambda-gt default 0.01 (the GT field is noisier than the teacher's; the old 0.1 made the GT anchor ~25% of the loss — set to 0 for pure distillation).
  • Text conditioning is the native FLUX.2-klein format: user-only + add_generation_prompt=True, enable_thinking=False, LEN=128 — the same rendered string the teacher's own tokenizer produces.
  • --compile / --attn-backend exist for A/B speed measurements (measure, do not assume). Measured on 2×RTX 5090 (2026-09-13, teacher-free): batch 24 + torch.compile ≈ 2.7-2.8 s/step vs 3.4-3.9 at batch 20 without compile (~+55% images/sec); the first ~11 steps after resume are slow (one compile per resolution bucket). No --cache-emb by preference (TE on GPU each step).

Samples

384×640 — teacher-free training, classic 20-pass model (~22k steps).

Ground truth (dataset latent → VAE decode):

GT 384x640

Student generation (same prompt, from noise):

Student 384x640

Held-out test — classic 20-pass, teacher-free ~21k steps (2026-09-14).

Girl

deer · fox

Status

  • Model: 1.996B, 20 passes (6 in + 8 mid + 6 out, 20 unique blocks), multi-layer text fusion (layers [2,12,22]).
  • Weights: classic — converted from the looped checkpoint; loop experiment dropped (second pass idled; classic == looped quality at ~14% less compute).
  • Training: teacher-free (train.py) at --batch-size 24 --compile, lr 5e-5, seed 43, on the merged 137k dataset (52k + testg 85k, auto-merged) — 2×RTX 5090.
Downloads last month
83
Safetensors
Model size
2B params
Tensor type
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for recoilme/sdxs

Finetuned
(65)
this model