diff --git a/README.es.md b/README.es.md index 7a313b804..4ec1c5d53 100644 --- a/README.es.md +++ b/README.es.md @@ -81,31 +81,51 @@ Para detalles de despliegue, consulta la [Guía Enterprise](/documentation/exper ### Compatibilidad de arquitectura de modelos -| Modelo | Parámetros | PEFT LoRA | Lycoris | Full-Rank | ControlNet | Cuantización | Flow Matching | Codificadores de texto | -|--------|------------|-----------|---------|-----------|------------|--------------|---------------|------------------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | int8 | ✗ | Ninguno | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | int8/nf4 (req.) | ✓ | T5-XXL | -| **SD 1.x/2.x (Legacy)** | 0.9B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L | - -*✓ = Compatible, ✗ = No compatible, * = Requiere DeepSpeed para entrenamiento full-rank* +SimpleTuner es compatible con las siguientes familias de modelos. El soporte detallado de funciones de entrenamiento está en la [Guía Quickstart](/documentation/QUICKSTART.es.md). + +| Modelo | Parámetros | Licencia | Uso comercial | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | Sí | +| **Anima** | No especificado | CircleStone Labs Non-Commercial License v1.2 | No (modelo); salidas permitidas | +| **Auraflow** | 6B | Apache-2.0 | Sí | +| **Boogu-Image** | No especificado | Apache-2.0 | Sí | +| **Chroma 1** | 8.9B | Apache-2.0 | Sí | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | Sí | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | Sí | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | No especificado | Apache-2.0 | Sí | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell); FLUX.1 [dev] Non-Commercial License (dev/Kontext) | Mixto por checkpoint | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B); FLUX Non-Commercial License (dev/klein 9B) | Mixto por checkpoint | +| **HeartMuLa** | 3B | No especificada en SimpleTuner | Ver términos upstream | +| **HiDream** | 17B (8.5B MoE) | MIT | Sí | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | Sí (copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | No | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | Sí | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | Sí | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | No especificado | Krea 2 Community License | Sí (menos de $1M en ingresos; salvaguardas requeridas) | +| **LongCat Image** | 6B | Apache-2.0 | Sí | +| **LongCat Video** | 13.6B | MIT | Sí | +| **LTX Video** | ~2.5B | Apache-2.0 | Sí | +| **LTX Video 2** | 19B | Apache-2.0 | Sí | +| **Lumina2** | 2B | Apache-2.0 | Sí | +| **Mage-Flow** | 4B | MIT | Sí | +| **OmniGen** | 3.8B | MIT | Sí | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | Sí (restringido) | +| **Qwen Image** | 20B | Apache-2.0 | Sí | +| **Sana** | 0.6B-4.8B | Apache-2.0 | Sí | +| **Sana Video** | 2B | Apache-2.0 | Sí | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | Sí (restringido) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | Sí (menos de $1M en ingresos) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | Sí (restringido) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | No especificada en SimpleTuner | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | Sí | +| **Wan S2V** | 14B | Apache-2.0 | Sí | +| **Z-Image** | 6B | Apache-2.0 | Sí | +| **Z-Image Omni** | 6B | Apache-2.0 | Sí | +| **ZLab I1** | 3B | MIT | Sí | + +*Los valores de licencia se toman de los helpers de modelos de SimpleTuner cuando están disponibles y de las tarjetas/licencias upstream para entradas que antes no estaban especificadas. `No especificada en SimpleTuner` significa que el helper no nombra una licencia y aquí no se resume ningún término upstream; revisa la tarjeta del modelo upstream antes de usarlo.* ### Técnicas avanzadas de entrenamiento diff --git a/README.hi.md b/README.hi.md index 77ff1b91c..699d5bb07 100644 --- a/README.hi.md +++ b/README.hi.md @@ -81,31 +81,51 @@ SimpleTuner एक पूर्ण मल्टी‑यूज़र प्र ### मॉडल आर्किटेक्चर समर्थन {#model-architecture-support} -| मॉडल | पैरामीटर | PEFT LoRA | Lycoris | फुल-रैंक | ControlNet | क्वांटाइज़ेशन | फ्लो मैचिंग | टेक्स्ट एन्कोडर | -|-------|------------|-----------|---------|-----------|------------|--------------|---------------|---------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | int8 | ✗ | कोई नहीं | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | int8/nf4 (req.) | ✓ | T5-XXL | -| **SD 1.x/2.x (Legacy)** | 0.9B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L | - -*✓ = समर्थित, ✗ = समर्थित नहीं, * = फुल‑रैंक प्रशिक्षण के लिए DeepSpeed आवश्यक* +SimpleTuner निम्नलिखित मॉडल families का समर्थन करता है। विस्तृत training feature समर्थन [Quickstart Guide](/documentation/QUICKSTART.hi.md) में है। + +| मॉडल | पैरामीटर | लाइसेंस | व्यावसायिक उपयोग | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | हाँ | +| **Anima** | निर्दिष्ट नहीं | CircleStone Labs Non-Commercial License v1.2 | नहीं (मॉडल); outputs अनुमत | +| **Auraflow** | 6B | Apache-2.0 | हाँ | +| **Boogu-Image** | निर्दिष्ट नहीं | Apache-2.0 | हाँ | +| **Chroma 1** | 8.9B | Apache-2.0 | हाँ | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | हाँ | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | हाँ | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | निर्दिष्ट नहीं | Apache-2.0 | हाँ | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell); FLUX.1 [dev] Non-Commercial License (dev/Kontext) | checkpoint के अनुसार mixed | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B); FLUX Non-Commercial License (dev/klein 9B) | checkpoint के अनुसार mixed | +| **HeartMuLa** | 3B | SimpleTuner में निर्दिष्ट नहीं | upstream शर्तें देखें | +| **HiDream** | 17B (8.5B MoE) | MIT | हाँ | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | हाँ (copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | नहीं | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | हाँ | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | हाँ | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | निर्दिष्ट नहीं | Krea 2 Community License | हाँ ($1M से कम राजस्व; safeguards आवश्यक) | +| **LongCat Image** | 6B | Apache-2.0 | हाँ | +| **LongCat Video** | 13.6B | MIT | हाँ | +| **LTX Video** | ~2.5B | Apache-2.0 | हाँ | +| **LTX Video 2** | 19B | Apache-2.0 | हाँ | +| **Lumina2** | 2B | Apache-2.0 | हाँ | +| **Mage-Flow** | 4B | MIT | हाँ | +| **OmniGen** | 3.8B | MIT | हाँ | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | हाँ (restricted) | +| **Qwen Image** | 20B | Apache-2.0 | हाँ | +| **Sana** | 0.6B-4.8B | Apache-2.0 | हाँ | +| **Sana Video** | 2B | Apache-2.0 | हाँ | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | हाँ (restricted) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | हाँ ($1M से कम राजस्व) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | हाँ (restricted) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | SimpleTuner में निर्दिष्ट नहीं | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | हाँ | +| **Wan S2V** | 14B | Apache-2.0 | हाँ | +| **Z-Image** | 6B | Apache-2.0 | हाँ | +| **Z-Image Omni** | 6B | Apache-2.0 | हाँ | +| **ZLab I1** | 3B | MIT | हाँ | + +*लाइसेंस मान उपलब्ध होने पर SimpleTuner model helpers से लिए गए हैं, और पहले अनिर्दिष्ट entries के लिए upstream model cards/licenses से। `SimpleTuner में निर्दिष्ट नहीं` का अर्थ है कि helper कोई license name नहीं देता और यहां upstream terms का सार नहीं दिया गया है; उपयोग से पहले upstream model card देखें.* ### उन्नत प्रशिक्षण तकनीकें {#advanced-training-techniques} diff --git a/README.ja.md b/README.ja.md index 4fa0b109c..c68739a3a 100644 --- a/README.ja.md +++ b/README.ja.md @@ -81,31 +81,51 @@ SimpleTunerには、エンタープライズグレードの機能を備えた完 ### モデルアーキテクチャサポート -| モデル | パラメータ数 | PEFT LoRA | Lycoris | Full-Rank | ControlNet | 量子化 | Flow Matching | テキストエンコーダー | -|-------|------------|-----------|---------|-----------|------------|--------------|---------------|---------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | int8 | ✗ | なし | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | int8/nf4 (req.) | ✓ | T5-XXL | -| **SD 1.x/2.x (Legacy)** | 0.9B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L | - -*✓ = サポート, ✗ = 非サポート, * = Full-rankトレーニングにDeepSpeedが必要* +SimpleTunerは以下のモデルファミリーをサポートしています。詳細なトレーニング機能の対応状況は[Quickstartガイド](/documentation/QUICKSTART.ja.md)を参照してください。 + +| モデル | パラメータ数 | ライセンス | 商用利用 | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | 可 | +| **Anima** | 未指定 | CircleStone Labs Non-Commercial License v1.2 | 不可(モデル);出力は可 | +| **Auraflow** | 6B | Apache-2.0 | 可 | +| **Boogu-Image** | 未指定 | Apache-2.0 | 可 | +| **Chroma 1** | 8.9B | Apache-2.0 | 可 | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | 可 | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | 可 | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | 未指定 | Apache-2.0 | 可 | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell); FLUX.1 [dev] Non-Commercial License (dev/Kontext) | checkpointごとに異なる | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B); FLUX Non-Commercial License (dev/klein 9B) | checkpointごとに異なる | +| **HeartMuLa** | 3B | SimpleTunerでは未指定 | 上流の条件を確認 | +| **HiDream** | 17B (8.5B MoE) | MIT | 可 | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | 可(copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | 不可 | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | 可 | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | 可 | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | 未指定 | Krea 2 Community License | 可(年収100万米ドル未満;安全対策必須) | +| **LongCat Image** | 6B | Apache-2.0 | 可 | +| **LongCat Video** | 13.6B | MIT | 可 | +| **LTX Video** | ~2.5B | Apache-2.0 | 可 | +| **LTX Video 2** | 19B | Apache-2.0 | 可 | +| **Lumina2** | 2B | Apache-2.0 | 可 | +| **Mage-Flow** | 4B | MIT | 可 | +| **OmniGen** | 3.8B | MIT | 可 | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | 可(制限あり) | +| **Qwen Image** | 20B | Apache-2.0 | 可 | +| **Sana** | 0.6B-4.8B | Apache-2.0 | 可 | +| **Sana Video** | 2B | Apache-2.0 | 可 | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | 可(制限あり) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | 可(年収100万米ドル未満) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | 可(制限あり) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | SimpleTunerでは未指定 | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | 可 | +| **Wan S2V** | 14B | Apache-2.0 | 可 | +| **Z-Image** | 6B | Apache-2.0 | 可 | +| **Z-Image Omni** | 6B | Apache-2.0 | 可 | +| **ZLab I1** | 3B | MIT | 可 | + +*ライセンス値は、利用可能な場合はSimpleTunerのモデルヘルパーから、以前未指定だった項目は上流のモデルカード/ライセンスから取得しています。`SimpleTunerでは未指定`は、ヘルパーがライセンス名を持たず、ここでも上流条件を要約していないことを意味します。使用前に上流モデルカードを確認してください。* ### 高度なトレーニング技術 diff --git a/README.md b/README.md index b38d7cf3c..3a0e88972 100644 --- a/README.md +++ b/README.md @@ -81,31 +81,51 @@ For deployment details, see the [Enterprise Guide](/documentation/experimental/s ### Model Architecture Support -| Model | Parameters | PEFT LoRA | Lycoris | Full-Rank | ControlNet | Ref Inputs | Quantization | Flow Matching | Text Encoders | -|-------|------------|-----------|---------|-----------|------------|------------|--------------|---------------|---------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | ✗ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | ✗ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | ✓ (Kontext) | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | ✓ opt | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | ✗ | int8 | ✗ | None | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | ✗ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | ✗ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | ✓ I2V | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | ✓ opt | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | ✗ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | int8/nf4 (req.) | ✓ | T5-XXL | -| **SD 1.x/2.x (Legacy)** | 0.9B | ✓ | ✓ | ✓ | ✓ | ✗ | int8/nf4 | ✗ | CLIP-L | - -*✓ = Supported, ✗ = Not supported, * = Requires DeepSpeed for full-rank training, Ref Inputs marks existing reference/edit/I2V conditioning paths only* +SimpleTuner supports the following model families. Detailed training feature support lives in the [Quickstart Guide](/documentation/QUICKSTART.md#feature-compatibility). + +| Model | Parameters | License | Commercial use | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | Yes | +| **Anima** | Not specified | CircleStone Labs Non-Commercial License v1.2 | No (model); outputs allowed | +| **Auraflow** | 6B | Apache-2.0 | Yes | +| **Boogu-Image** | Not specified | Apache-2.0 | Yes | +| **Chroma 1** | 8.9B | Apache-2.0 | Yes | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | Yes | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | Yes | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | Not specified | Apache-2.0 | Yes | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell); FLUX.1 [dev] Non-Commercial License (dev/Kontext) | Mixed by checkpoint | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B); FLUX Non-Commercial License (dev/klein 9B) | Mixed by checkpoint | +| **HeartMuLa** | 3B | Not specified in SimpleTuner | See upstream terms | +| **HiDream** | 17B (8.5B MoE) | MIT | Yes | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | Yes (copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | No | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | Yes | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | Yes | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | Not specified | Krea 2 Community License | Yes (under $1M revenue; safeguards required) | +| **LongCat Image** | 6B | Apache-2.0 | Yes | +| **LongCat Video** | 13.6B | MIT | Yes | +| **LTX Video** | ~2.5B | Apache-2.0 | Yes | +| **LTX Video 2** | 19B | Apache-2.0 | Yes | +| **Lumina2** | 2B | Apache-2.0 | Yes | +| **Mage-Flow** | 4B | MIT | Yes | +| **OmniGen** | 3.8B | MIT | Yes | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | Yes (restricted) | +| **Qwen Image** | 20B | Apache-2.0 | Yes | +| **Sana** | 0.6B-4.8B | Apache-2.0 | Yes | +| **Sana Video** | 2B | Apache-2.0 | Yes | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | Yes (restricted) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | Yes (under $1M revenue) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | Yes (restricted) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | Not specified in SimpleTuner | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | Yes | +| **Wan S2V** | 14B | Apache-2.0 | Yes | +| **Z-Image** | 6B | Apache-2.0 | Yes | +| **Z-Image Omni** | 6B | Apache-2.0 | Yes | +| **ZLab I1** | 3B | MIT | Yes | + +*License values are taken from SimpleTuner model helpers when available and upstream model cards/licenses for entries that were previously unspecified. `Not specified in SimpleTuner` means the helper does not name a license and no upstream term is summarized here; check the upstream model card before use.* ### Advanced Training Techniques diff --git a/README.pt-BR.md b/README.pt-BR.md index 19b7ca412..a4ccf22aa 100644 --- a/README.pt-BR.md +++ b/README.pt-BR.md @@ -81,31 +81,51 @@ Para detalhes de deploy, veja o [guia enterprise](/documentation/experimental/se ### Suporte a arquitetura de modelos -| Modelo | Parametros | PEFT LoRA | Lycoris | Full-Rank | ControlNet | Quantizacao | Flow Matching | Text Encoders | -|-------|------------|-----------|---------|-----------|------------|--------------|---------------|---------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | int8 | ✗ | Nenhum | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | int8/nf4 (req.) | ✓ | T5-XXL | -| **SD 1.x/2.x (Legacy)** | 0.9B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L | - -*✓ = Suportado, ✗ = Nao suportado, * = Requer DeepSpeed para treino full-rank* +SimpleTuner suporta as seguintes familias de modelos. O suporte detalhado a recursos de treinamento fica no [guia Quickstart](/documentation/QUICKSTART.pt-BR.md). + +| Modelo | Parametros | Licenca | Uso comercial | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | Sim | +| **Anima** | Nao especificado | CircleStone Labs Non-Commercial License v1.2 | Nao (modelo); outputs permitidos | +| **Auraflow** | 6B | Apache-2.0 | Sim | +| **Boogu-Image** | Nao especificado | Apache-2.0 | Sim | +| **Chroma 1** | 8.9B | Apache-2.0 | Sim | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | Sim | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | Sim | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | Nao especificado | Apache-2.0 | Sim | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell); FLUX.1 [dev] Non-Commercial License (dev/Kontext) | Misto por checkpoint | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B); FLUX Non-Commercial License (dev/klein 9B) | Misto por checkpoint | +| **HeartMuLa** | 3B | Nao especificada no SimpleTuner | Ver termos upstream | +| **HiDream** | 17B (8.5B MoE) | MIT | Sim | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | Sim (copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | Nao | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | Sim | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | Sim | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | Nao especificado | Krea 2 Community License | Sim (menos de $1M em receita; salvaguardas exigidas) | +| **LongCat Image** | 6B | Apache-2.0 | Sim | +| **LongCat Video** | 13.6B | MIT | Sim | +| **LTX Video** | ~2.5B | Apache-2.0 | Sim | +| **LTX Video 2** | 19B | Apache-2.0 | Sim | +| **Lumina2** | 2B | Apache-2.0 | Sim | +| **Mage-Flow** | 4B | MIT | Sim | +| **OmniGen** | 3.8B | MIT | Sim | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | Sim (restrito) | +| **Qwen Image** | 20B | Apache-2.0 | Sim | +| **Sana** | 0.6B-4.8B | Apache-2.0 | Sim | +| **Sana Video** | 2B | Apache-2.0 | Sim | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | Sim (restrito) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | Sim (menos de $1M em receita) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | Sim (restrito) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | Nao especificada no SimpleTuner | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | Sim | +| **Wan S2V** | 14B | Apache-2.0 | Sim | +| **Z-Image** | 6B | Apache-2.0 | Sim | +| **Z-Image Omni** | 6B | Apache-2.0 | Sim | +| **ZLab I1** | 3B | MIT | Sim | + +*Os valores de licenca vem dos model helpers do SimpleTuner quando disponiveis e dos model cards/licencas upstream para entradas que antes nao estavam especificadas. `Nao especificada no SimpleTuner` significa que o helper nao informa uma licenca e nenhum termo upstream esta resumido aqui; verifique o model card upstream antes de usar.* ### Tecnicas avancadas de treinamento diff --git a/README.zh.md b/README.zh.md index 7e4995b63..e1ce18596 100644 --- a/README.zh.md +++ b/README.zh.md @@ -81,31 +81,51 @@ SimpleTuner 包含完整的多用户训练平台,具有企业级功能——** ### 模型架构支持 -| 模型 | 参数量 | PEFT LoRA | Lycoris | 全秩 | ControlNet | 量化 | Flow Matching | 文本编码器 | -|-------|------------|-----------|---------|-----------|------------|--------------|---------------|---------------| -| **Stable Diffusion XL** | 3.5B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L/G | -| **Stable Diffusion 3** | 2B-8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L/G + T5-XXL | -| **Flux.1** | 12B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL | -| **Flux.2** | 32B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | Mistral-3 Small | -| **Ideogram 4** | 9B | ✓ | ✓ | ✓* | ✗ | fp8/nf4 | ✓ | Qwen3-VL | -| **ACE-Step** | 3.5B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HeartMuLa** | 3B | ✓ | ✓ | ✓* | ✗ | int8 | ✗ | 无 | -| **Chroma 1** | 8.9B | ✓ | ✓ | ✓* | ✗ | int8/fp8/nf4 | ✓ | T5-XXL | -| **Auraflow** | 6.8B | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | UMT5-XXL | -| **PixArt Sigma** | 0.6B-0.9B | ✗ | ✓ | ✓ | ✓ | int8 | ✗ | T5-XXL | -| **Sana** | 0.6B-4.8B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2-2B | -| **Lumina2** | 2B | ✓ | ✓ | ✓ | ✗ | int8 | ✓ | Gemma2 | -| **Kwai Kolors** | 5B | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ChatGLM-6B | -| **LTX Video** | 5B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **LTX Video 2** | 19B | ✓ | ✓ | ✓* | ✗ | int8/fp8 | ✓ | Gemma3 | -| **Wan Video** | 1.3B-14B | ✓ | ✓ | ✓* | ✗ | int8 | ✓ | UMT5 | -| **HiDream** | 17B (8.5B MoE) | ✓ | ✓ | ✓* | ✓ | int8/fp8/nf4 | ✓ | CLIP-L + T5-XXL + Llama | -| **Cosmos2** | 2B-14B | ✗ | ✓ | ✓ | ✗ | int8 | ✓ | T5-XXL | -| **OmniGen** | 3.8B | ✓ | ✓ | ✓ | ✗ | int8/fp8 | ✓ | T5-XXL | -| **Qwen Image** | 20B | ✓ | ✓ | ✓* | ✗ | int8/nf4(必需) | ✓ | T5-XXL | -| **SD 1.x/2.x(旧版)** | 0.9B | ✓ | ✓ | ✓ | ✓ | int8/nf4 | ✗ | CLIP-L | - -*✓ = 支持,✗ = 不支持,* = 全秩训练需要 DeepSpeed* +SimpleTuner 支持以下模型系列。详细的训练功能支持请参阅[快速入门指南](/documentation/QUICKSTART.zh.md)。 + +| 模型 | 参数量 | 许可证 | 商业使用 | +| --- | --- | --- | --- | +| **ACE-Step** | 3.5B | Apache-2.0 | 是 | +| **Anima** | 未指定 | CircleStone Labs Non-Commercial License v1.2 | 否(模型);输出可商用 | +| **Auraflow** | 6B | Apache-2.0 | 是 | +| **Boogu-Image** | 未指定 | Apache-2.0 | 是 | +| **Chroma 1** | 8.9B | Apache-2.0 | 是 | +| **Cosmos2** | 2B-14B | NVIDIA Open Model License | 是 | +| **Cosmos3** | 16B-65B | OpenMDW-1.1 | 是 | +| **DeepFloyd IF** | 0.4B-4.3B stages | DeepFloyd IF License | Abandonware | +| **ERNIE-Image** | 未指定 | Apache-2.0 | 是 | +| **Flux.1** | 8B-12B | Apache-2.0 (schnell);FLUX.1 [dev] Non-Commercial License (dev/Kontext) | 按 checkpoint 不同 | +| **Flux.2** | 4B-32B | Apache-2.0 (klein 4B);FLUX Non-Commercial License (dev/klein 9B) | 按 checkpoint 不同 | +| **HeartMuLa** | 3B | SimpleTuner 中未指定 | 查看上游条款 | +| **HiDream** | 17B (8.5B MoE) | MIT | 是 | +| **Hunyuan Video** | 8.3B | AGPL-3.0 | 是(copyleft) | +| **Ideogram 4** | 9B | Ideogram 4 Non-Commercial | 否 | +| **Kandinsky 5.0 Image** | 6B (lite) | MIT | 是 | +| **Kandinsky 5.0 Video** | 2B lite, 19B pro | MIT | 是 | +| **Kwai Kolors** | 2.7B | Apache-2.0 | Abandonware | +| **Krea2** | 未指定 | Krea 2 Community License | 是(年收入低于 100 万美元;需安全防护) | +| **LongCat Image** | 6B | Apache-2.0 | 是 | +| **LongCat Video** | 13.6B | MIT | 是 | +| **LTX Video** | ~2.5B | Apache-2.0 | 是 | +| **LTX Video 2** | 19B | Apache-2.0 | 是 | +| **Lumina2** | 2B | Apache-2.0 | 是 | +| **Mage-Flow** | 4B | MIT | 是 | +| **OmniGen** | 3.8B | MIT | 是 | +| **PixArt Sigma** | 0.6B-0.9B | OpenRAIL++ | 是(受限) | +| **Qwen Image** | 20B | Apache-2.0 | 是 | +| **Sana** | 0.6B-4.8B | Apache-2.0 | 是 | +| **Sana Video** | 2B | Apache-2.0 | 是 | +| **SD 1.x/2.x (Legacy)** | 0.9B | OpenRAIL++ | 是(受限) | +| **Stable Diffusion 3** | 2B-8B | Stability AI Community License | 是(年收入低于 100 万美元) | +| **Stable Diffusion XL** | 3.5B | CreativeML OpenRAIL-M | 是(受限) | +| **Stable Cascade (Stage C)** | 1B, 3.6B prior | SimpleTuner 中未指定 | Abandonware | +| **Wan Video** | 1.3B-14B | Apache-2.0 | 是 | +| **Wan S2V** | 14B | Apache-2.0 | 是 | +| **Z-Image** | 6B | Apache-2.0 | 是 | +| **Z-Image Omni** | 6B | Apache-2.0 | 是 | +| **ZLab I1** | 3B | MIT | 是 | + +*许可证值在可用时来自 SimpleTuner 模型 helper,并为此前未指定的条目参考上游模型卡/许可证。`SimpleTuner 中未指定` 表示 helper 未命名许可证,且此处未汇总上游条款;使用前请查看上游模型卡。* ### 高级训练技术 diff --git a/documentation/OPTIONS.es.md b/documentation/OPTIONS.es.md index 961a91527..61bb37a60 100644 --- a/documentation/OPTIONS.es.md +++ b/documentation/OPTIONS.es.md @@ -253,6 +253,13 @@ Donde `foo` es tu entorno de configuración; o simplemente usa `config/config.js - **Qué**: Ruta al modelo Gemma preentrenado o su identificador en . - **Por qué**: Al entrenar modelos basados en Gemma (por ejemplo LTX-2, Sana o Lumina2), puedes apuntar a un checkpoint Gemma compartido sin cambiar la ruta del modelo base de difusión. +### `--qwen_text_encoder_model_name_or_path` + +- **Qué**: Ruta a un codificador de texto Qwen preentrenado o su identificador en . +- **Predeterminado**: `None` (usa la fuente del codificador de texto Qwen definida por el modelo seleccionado). +- **Por qué**: Úsalo para compartir o reemplazar el codificador de texto Qwen en familias de modelos basadas en Qwen sin editar la caché de Hugging Face. +- **Notas**: Se aplica a familias de modelos con un solo codificador de texto Qwen. Si una familia define varios codificadores Qwen, la opción se ignora y SimpleTuner registra una advertencia. + ### `--max_grounding_entities` - Numero maximo de entidades de grounding por imagen para anotaciones espaciales estilo GLIGEN. Por defecto: 0 (deshabilitado). Valores tipicos: 4-16. @@ -275,8 +282,14 @@ Donde `foo` es tu entorno de configuración; o simplemente usa `config/config.js ### `--gradient_checkpointing_interval` -- **Qué**: Hace checkpoint cada *n* bloques, donde *n* es un valor mayor que cero. Un valor de 1 es equivalente a dejar `--gradient_checkpointing` habilitado, y un valor de 2 hará checkpoint en bloques alternos. -- **Nota**: SDXL y Flux son actualmente los únicos modelos que soportan esta opción. SDXL usa una implementación algo improvisada. +- **Qué**: Intervalo dependiente del modelo para checkpointing de bloques transformer. Un valor de 1 equivale básicamente a dejar `--gradient_checkpointing` habilitado. +- **Nota**: Flux, Flux.2, Krea 2, LTXVideo2, MageFlow, Z-Image y Wan usan chunks contiguos de *n* bloques en rutas whole-block. Otras familias que exponen esta opción pueden seguir usando el comportamiento anterior de "checkpoint cada *n* bloques". Valores más altos pueden reducir recompute, pero normalmente dejan más activaciones en VRAM. + +### `--gradient_checkpointing_segment_stride` + +- **Qué**: Inicia un segmento con checkpoint cada *n* bloques en rutas segmented whole-block compatibles. +- **Ejemplo**: Con `--gradient_checkpointing_interval=2` y `--gradient_checkpointing_segment_stride=4`, SimpleTuner checkpointa dos bloques, ejecuta los dos siguientes normalmente y repite. +- **Nota**: Solo tiene efecto en familias de modelos que exponen soporte segmented whole-block en la version instalada de SimpleTuner. Las familias no soportadas registran una advertencia e ignoran el valor. El stride debe ser al menos igual al interval. Consulta [Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md). ### `--gradient_checkpointing_backend` @@ -286,7 +299,27 @@ Donde `foo` es tu entorno de configuración; o simplemente usa `config/config.js - `torch-ffn`: checkpoint solo del lado feed-forward en modelos con un límite FFN limpio. - `unsloth`: checkpoint del bloque completo compatible y offload de tensores guardados a CPU. - `unsloth-ffn`: checkpoint solo del lado feed-forward y offload de sus tensores guardados a CPU. -- **Nota**: Solo efectivo cuando `--gradient_checkpointing` está habilitado. Las variantes `unsloth` requieren CUDA. Las variantes FFN-only soportan actualmente bloques estilo Flux.1 y MageFlow, y fallan de forma explícita si no existe ese scope. Consulta [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) para tradeoffs medidos. +- **Nota**: Solo efectivo cuando `--gradient_checkpointing` está habilitado. Las variantes `unsloth` requieren CUDA. Las variantes FFN-only soportan actualmente Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan y Z-Image, y fallan de forma explícita si no existe ese scope. Consulta [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) para tradeoffs medidos. + +### `--gradient_checkpointing_offload_attention` + +- **Qué**: Hace offload a CPU de las activaciones guardadas del lado attention en modelos con un límite attention/FFN limpio. +- **Por qué**: Cuando transferir es más barato que recomputar attention, reduce VRAM sin pagar todo el coste de rematerializar attention. +- **Nota**: Puede activarse por si solo. Tambien puede combinarse con cualquier checkpoint backend que soporte el modelo. Solo tiene efecto en familias de modelos que exponen una frontera attention/FFN limpia en la version instalada de SimpleTuner; las familias no soportadas fallan de forma explicita. + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **Predeterminado**: `12` +- **Qué**: Número máximo de buckets distintos de tensores CPU pinned usados por activation offload. +- **Por qué**: Pinned memory mejora las transferencias CPU/GPU, pero resoluciones y longitudes de texto variables pueden crear formas raras. Al alcanzar este límite, las nuevas formas de bucket usan memoria CPU normal. +- **Nota**: Usa `0` para desactivar el pooling de pinned memory para activation offload. + +### `--gradient_checkpointing_offload_prefetch` + +- **Predeterminado**: `false` +- **Qué**: Aprende el orden de restore en backward para activations offloaded y precarga en GPU el tensor que probablemente venga después. +- **Por qué**: El restore H2D justo a tiempo casi no se puede solapar. Con un orden estable, prefetch puede ocultar parte de la transferencia detrás del backward compute. +- **Nota**: Experimental y solo activo con `--gradient_checkpointing_offload_attention`. ### `--refiner_training` @@ -1734,6 +1767,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1748,6 +1782,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2050,6 +2085,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2079,6 +2116,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.hi.md b/documentation/OPTIONS.hi.md index 448ebaf9c..c9377a630 100644 --- a/documentation/OPTIONS.hi.md +++ b/documentation/OPTIONS.hi.md @@ -253,6 +253,13 @@ simpletuner configure config/foo/config.json - **What**: pretrained Gemma model का path या से उसका identifier. - **Why**: Gemma‑based models (जैसे LTX-2, Sana, Lumina2) ट्रेन करते समय आप base diffusion model path बदले बिना Gemma weights का source specify कर सकते हैं। +### `--qwen_text_encoder_model_name_or_path` + +- **What**: pretrained Qwen text encoder model का path या से उसका identifier. +- **Default**: `None` (selected model में defined Qwen text encoder source उपयोग होता है). +- **Why**: Qwen-based model families में Qwen text encoder को share या replace करने के लिए इसका उपयोग करें, बिना Hugging Face cache edit किए। +- **Notes**: यह उन model families पर लागू होता है जिनमें एक Qwen text encoder है। अगर कोई model family multiple Qwen encoders define करती है, तो option ignore होता है और SimpleTuner warning log करता है। + ### `--max_grounding_entities` - GLIGEN-style spatial annotations के लिए प्रति image grounding entities की अधिकतम संख्या। Default: 0 (disabled)। सामान्य मान: 4-16। @@ -275,8 +282,14 @@ simpletuner configure config/foo/config.json ### `--gradient_checkpointing_interval` -- **What**: हर *n* blocks पर checkpoint करें, जहाँ *n* शून्य से बड़ा मान है। 1 का मान `--gradient_checkpointing` enabled जैसा है, और 2 हर दूसरे block पर checkpoint करेगा। -- **Note**: यह विकल्प फिलहाल केवल SDXL और Flux में समर्थित है। SDXL इसमें hackish implementation उपयोग करता है। +- **What**: Transformer block checkpointing के लिए model-dependent interval। 1 का मान लगभग `--gradient_checkpointing` enabled जैसा है। +- **Note**: Flux, Flux.2, Krea 2, LTXVideo2, MageFlow, Z-Image, और Wan whole-block paths पर *n* contiguous block chunks use करते हैं। इस option को expose करने वाली दूसरी families अभी भी पुराने "हर *n*-th block checkpoint" behavior का उपयोग कर सकती हैं। Higher values recompute overhead घटा सकती हैं, लेकिन आम तौर पर VRAM में ज्यादा activations रखती हैं। + +### `--gradient_checkpointing_segment_stride` + +- **What**: Supported segmented whole-block paths पर हर *n* blocks में checkpointed segment शुरू करें। +- **Example**: `--gradient_checkpointing_interval=2` और `--gradient_checkpointing_segment_stride=4` के साथ SimpleTuner दो blocks checkpoint करता है, अगले दो blocks normally चलाता है, और repeat करता है। +- **Note**: यह केवल उन model families पर प्रभावी होता है जो installed SimpleTuner version में segmented whole-block support expose करती हैं। Unsupported families warning log करती हैं और value ignore करती हैं। stride interval से कम नहीं हो सकता। [Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) देखें। ### `--gradient_checkpointing_backend` @@ -286,7 +299,27 @@ simpletuner configure config/foo/config.json - `torch-ffn`: साफ FFN boundary वाले models पर केवल feed-forward side checkpoint करता है। - `unsloth`: पूरे supported block को checkpoint करता है और saved tensors CPU पर offload करता है। - `unsloth-ffn`: केवल feed-forward side checkpoint करता है और उसके saved tensors CPU पर offload करता है। -- **Note**: केवल `--gradient_checkpointing` enabled होने पर प्रभावी। `unsloth` variants के लिए CUDA आवश्यक है। FFN-only variants अभी Flux.1-style blocks और MageFlow support करते हैं, और unsupported scope पर साफ error मिलता है। Measured tradeoffs के लिए [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) देखें। +- **Note**: केवल `--gradient_checkpointing` enabled होने पर प्रभावी। `unsloth` variants के लिए CUDA आवश्यक है। FFN-only variants अभी Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan, और Z-Image support करते हैं, और unsupported scope पर साफ error मिलता है। Measured tradeoffs के लिए [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) देखें। + +### `--gradient_checkpointing_offload_attention` + +- **What**: साफ attention/FFN boundary वाले models पर attention-side saved activations CPU पर offload करें। +- **Why**: जब transfer attention recompute से सस्ता हो, तो full attention rematerialization cost दिए बिना VRAM घट सकती है। +- **Note**: इसे अकेले enable किया जा सकता है। Model जिस भी checkpoint backend को support करे, उसके साथ भी combine किया जा सकता है। यह केवल उन model families पर प्रभावी होता है जो installed SimpleTuner version में साफ attention/FFN boundary expose करती हैं; unsupported model families साफ error देती हैं। + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **Default**: `12` +- **What**: activation offload द्वारा इस्तेमाल किए जाने वाले अलग-अलग pinned CPU tensor buckets की maximum संख्या। +- **Why**: pinned memory CPU/GPU transfer में मदद करती है, लेकिन variable resolution और text length rare tensor shapes बना सकते हैं। limit पूरी होने के बाद नए bucket shapes normal CPU memory use करेंगे। +- **Note**: activation offload के pinned-memory pooling को बंद करने के लिए `0` set करें। + +### `--gradient_checkpointing_offload_prefetch` + +- **Default**: `false` +- **What**: offloaded activations के backward restore order को learn करके likely next tensor को GPU पर prefetch करता है। +- **Why**: JIT H2D restore अक्सर overlap नहीं कर पाता। order stable होने के बाद prefetch transfer के कुछ हिस्से को backward compute के पीछे छिपा सकता है। +- **Note**: Experimental है और केवल `--gradient_checkpointing_offload_attention` के साथ active होता है। ### `--refiner_training` @@ -1732,6 +1765,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1746,6 +1780,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2048,6 +2083,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2077,6 +2114,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.ja.md b/documentation/OPTIONS.ja.md index 49b6b2b3a..d3f435367 100644 --- a/documentation/OPTIONS.ja.md +++ b/documentation/OPTIONS.ja.md @@ -254,6 +254,13 @@ simpletuner configure config/foo/config.json - **内容**: 事前学習済み Gemma モデルのパス、または の識別子。 - **理由**: Gemma 系モデル(例: LTX-2、Sana、Lumina2)を学習する際、ベース拡散モデルのパスを変えずに Gemma 重みの参照先を指定できます。 +### `--qwen_text_encoder_model_name_or_path` + +- **内容**: 事前学習済み Qwen テキストエンコーダーモデルのパス、または の識別子。 +- **既定**: `None`(選択したモデルが定義する Qwen テキストエンコーダーの参照元を使用します) +- **理由**: Hugging Face キャッシュを編集せずに、Qwen 系モデルファミリーの Qwen テキストエンコーダーを共有または置き換えるために使用します。 +- **注記**: Qwen テキストエンコーダーが 1 つのモデルファミリーに適用されます。複数の Qwen エンコーダーを定義するモデルファミリーでは、このオプションは無視され、SimpleTuner が警告を記録します。 + ### `--max_grounding_entities` - GLIGEN スタイルの空間アノテーション用に、画像あたりのグラウンディングエンティティの最大数を指定します。デフォルト: 0(無効)。一般的な値: 4-16。 @@ -276,8 +283,14 @@ simpletuner configure config/foo/config.json ### `--gradient_checkpointing_interval` -- **内容**: *n* ブロックごとにチェックポイントを作成します。値は 0 より大きい必要があります。1 は `--gradient_checkpointing` と同等で、2 は隔ブロックでチェックポイントを作成します。 -- **注記**: 現在このオプションに対応しているのは SDXL と Flux のみです。SDXL は暫定的な実装です。 +- **内容**: transformer block checkpointing のモデル依存 interval です。1 は `--gradient_checkpointing` を有効にした状態とほぼ同じです。 +- **注記**: Flux、Flux.2、Krea 2、LTXVideo2、MageFlow、Z-Image、Wan は whole-block path で連続した *n* block chunk を使います。このオプションを持つ他の family は、従来の「*n* block ごとに checkpoint」挙動のままの場合があります。値を大きくすると再計算 overhead は減ることがありますが、通常は VRAM に残る activation が増えます。 + +### `--gradient_checkpointing_segment_stride` + +- **内容**: 対応する segmented whole-block path で、*n* block ごとに checkpointed segment を開始します。 +- **例**: `--gradient_checkpointing_interval=2` と `--gradient_checkpointing_segment_stride=4` では、SimpleTuner は 2 block を checkpoint し、次の 2 block を通常実行し、それを繰り返します。 +- **注記**: インストール済み SimpleTuner のバージョンで segmented whole-block support を公開している model family でのみ有効です。未対応 family では warning を記録し、値を無視します。stride は interval 以上である必要があります。[Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) を参照してください。 ### `--gradient_checkpointing_backend` @@ -287,7 +300,27 @@ simpletuner configure config/foo/config.json - `torch-ffn`: 明確な FFN 境界があるモデルで feed-forward 側だけを checkpoint します。 - `unsloth`: 対応 block 全体を checkpoint し、保存 tensor を CPU に offload します。 - `unsloth-ffn`: feed-forward 側だけを checkpoint し、保存 tensor を CPU に offload します。 -- **注記**: `--gradient_checkpointing` が有効な場合のみ機能します。`unsloth` 系は CUDA が必要です。FFN-only 系は現在 Flux.1-style blocks と MageFlow に対応し、対応していない scope では明示的に失敗します。実測 tradeoff は [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) を参照してください。 +- **注記**: `--gradient_checkpointing` が有効な場合のみ機能します。`unsloth` 系は CUDA が必要です。FFN-only 系は現在 Chroma、Flux、Krea 2、LTXVideo2、MageFlow、Wan、Z-Image に対応し、対応していない scope では明示的に失敗します。実測 tradeoff は [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) を参照してください。 + +### `--gradient_checkpointing_offload_attention` + +- **内容**: 明確な attention/FFN 境界があるモデルで、attention 側の保存 activations を CPU に offload します。 +- **理由**: attention を再計算するより転送が安い場合、完全な attention rematerialization cost を払わずに VRAM を減らせます。 +- **注記**: 単体でも有効化できます。そのモデルがサポートする任意の checkpoint backend とも組み合わせられます。インストール済み SimpleTuner のバージョンで明確な attention/FFN boundary を公開している model family でのみ有効です。未対応モデルでは明示的に失敗します。 + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **デフォルト**: `12` +- **内容**: activation offload が使う pinned CPU tensor bucket の最大数です。 +- **理由**: pinned memory は CPU/GPU 転送に有利ですが、可変解像度や可変 text length では珍しい tensor shape が出ます。上限に達した後の新しい bucket shape は通常の CPU memory を使います。 +- **注記**: `0` にすると activation offload の pinned-memory pooling を無効化します。 + +### `--gradient_checkpointing_offload_prefetch` + +- **デフォルト**: `false` +- **内容**: offload された activations の backward restore order を学習し、次に必要になりそうな tensor を GPU に prefetch します。 +- **理由**: JIT H2D restore はほとんど overlap できません。順序が安定すると、prefetch は一部の転送を backward compute の裏に隠せます。 +- **注記**: 実験的機能で、`--gradient_checkpointing_offload_attention` が有効な場合のみ動作します。 ### `--refiner_training` @@ -1735,6 +1768,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1749,6 +1783,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2050,6 +2085,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2079,6 +2116,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.md b/documentation/OPTIONS.md index ebd0a1170..71308a0f3 100644 --- a/documentation/OPTIONS.md +++ b/documentation/OPTIONS.md @@ -253,6 +253,13 @@ Where `foo` is your config environment - or just use `config/config.json` if you - **What**: Path to the pretrained Gemma model or its identifier from . - **Why**: When training Gemma-based models (for example LTX-2, Sana, or Lumina2), you can point at a shared Gemma checkpoint without changing the base diffusion model path. +### `--qwen_text_encoder_model_name_or_path` + +- **What**: Path to a pretrained Qwen text encoder model or its identifier from . +- **Default**: `None` (use the Qwen text encoder source defined by the selected model). +- **Why**: Use this to share or replace the Qwen text encoder used by Qwen-based model families without editing the Hugging Face cache. +- **Notes**: This applies to model families with one Qwen text encoder. If a model family defines multiple Qwen text encoders, the option is ignored and SimpleTuner logs a warning. + ### `--max_grounding_entities` - **What**: Maximum number of grounding entities per image for GLIGEN-style spatial annotations. @@ -282,8 +289,14 @@ Where `foo` is your config environment - or just use `config/config.json` if you ### `--gradient_checkpointing_interval` -- **What**: Checkpoint only every *n* blocks, where *n* is a value greater than zero. A value of 1 is effectively the same as just leaving `--gradient_checkpointing` enabled, and a value of 2 will checkpoint every other block. -- **Note**: SDXL and Flux are currently the only models supporting this option. SDXL uses a hackish implementation. +- **What**: Model-dependent interval for transformer block checkpointing. A value of 1 is effectively the same as leaving `--gradient_checkpointing` enabled. +- **Note**: Flux, Flux.2, Krea 2, LTXVideo2, MageFlow, Z-Image, and Wan use contiguous chunks of *n* blocks on whole-block paths. Other families that expose this option may use the older "checkpoint every *n*-th block" behavior. Higher values can reduce recompute overhead, but usually keep more activations in VRAM. + +### `--gradient_checkpointing_segment_stride` + +- **What**: Start a checkpointed segment every *n* blocks on supported segmented whole-block paths. +- **Example**: With `--gradient_checkpointing_interval=2` and `--gradient_checkpointing_segment_stride=4`, SimpleTuner checkpoints two blocks, runs the next two blocks normally, and repeats. +- **Note**: Only takes effect on model families that expose segmented whole-block support in the installed SimpleTuner version. Unsupported families log a warning and ignore the value. The stride must be at least the interval. See [Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md). ### `--gradient_checkpointing_backend` @@ -293,7 +306,27 @@ Where `foo` is your config environment - or just use `config/config.json` if you - `torch-ffn`: checkpoint only the feed-forward side on models that expose a clean FFN boundary. - `unsloth`: checkpoint the whole supported block and offload saved tensors to CPU. - `unsloth-ffn`: checkpoint only the feed-forward side and offload its saved tensors to CPU. -- **Note**: Only effective when `--gradient_checkpointing` is enabled. The `unsloth` variants require CUDA. FFN-only variants currently support Flux.1-style blocks and MageFlow, and fail loudly when the model does not expose that scope. See [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) for measured tradeoffs. +- **Note**: Only effective when `--gradient_checkpointing` is enabled. The `unsloth` variants require CUDA. FFN-only variants currently support Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan, and Z-Image, and fail loudly when the model does not expose that scope. See [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) for measured tradeoffs. + +### `--gradient_checkpointing_offload_attention` + +- **What**: Offload attention-side saved activations to CPU on models with a clean attention/FFN boundary. +- **Why**: When transfer is cheaper than recomputing attention, this can reduce VRAM without paying the full attention rematerialization cost. +- **Note**: This can be enabled by itself. It can also be combined with any checkpoint backend that the model supports. It only takes effect on model families that expose a clean attention/FFN boundary in the installed SimpleTuner version; unsupported model families fail loudly. + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **Default**: `12` +- **What**: Maximum number of distinct pinned CPU tensor buckets used by activation offload. +- **Why**: Pinned memory improves CPU/GPU transfer behavior, but variable resolutions and text lengths can create rare tensor shapes. Once this cap is reached, new bucket shapes use normal CPU memory instead. +- **Note**: Set `0` to disable pinned-memory pooling for activation offload. + +### `--gradient_checkpointing_offload_prefetch` + +- **Default**: `false` +- **What**: Learn the backward restore order for labeled offloaded activations and prefetch likely next tensors back to GPU. +- **Why**: JIT H2D restore usually cannot overlap much. Prefetch can hide some transfer behind backward compute after the order stabilizes. +- **Note**: Experimental and only active with `--gradient_checkpointing_offload_attention`. ### `--refiner_training` @@ -1738,6 +1771,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1752,6 +1786,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2054,6 +2089,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2083,6 +2120,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.pt-BR.md b/documentation/OPTIONS.pt-BR.md index 735d986e5..10c5e132a 100644 --- a/documentation/OPTIONS.pt-BR.md +++ b/documentation/OPTIONS.pt-BR.md @@ -253,6 +253,13 @@ Onde `foo` e seu ambiente de config — ou use `config/config.json` se nao estiv - **O que**: Caminho para o modelo Gemma pre-treinado ou seu identificador em . - **Por que**: Ao treinar modelos baseados em Gemma (por exemplo LTX-2, Sana ou Lumina2), voce pode apontar para um checkpoint Gemma compartilhado sem mudar o caminho do modelo base de difusao. +### `--qwen_text_encoder_model_name_or_path` + +- **O que**: Caminho para um encoder de texto Qwen pre-treinado ou seu identificador em . +- **Padrao**: `None` (usa a origem do encoder de texto Qwen definida pelo modelo selecionado). +- **Por que**: Use para compartilhar ou substituir o encoder de texto Qwen em familias de modelos baseadas em Qwen sem editar o cache do Hugging Face. +- **Notas**: Aplica-se a familias de modelos com um unico encoder de texto Qwen. Se uma familia definir varios encoders Qwen, a opcao e ignorada e o SimpleTuner registra um aviso. + ### `--max_grounding_entities` - Numero maximo de entidades de grounding por imagem para anotacoes espaciais no estilo GLIGEN. Padrao: 0 (desabilitado). Valores tipicos: 4-16. @@ -275,8 +282,14 @@ Onde `foo` e seu ambiente de config — ou use `config/config.json` se nao estiv ### `--gradient_checkpointing_interval` -- **O que**: Faz checkpoint apenas a cada *n* blocos, onde *n* e um valor maior que zero. Um valor 1 e efetivamente o mesmo que deixar `--gradient_checkpointing` habilitado, e 2 faz checkpoint a cada outro bloco. -- **Nota**: SDXL e Flux sao atualmente os unicos modelos que suportam essa opcao. SDXL usa uma implementacao meio hack. +- **O que**: Intervalo dependente do modelo para checkpointing de blocos transformer. Um valor 1 e basicamente o mesmo que deixar `--gradient_checkpointing` habilitado. +- **Nota**: Flux, Flux.2, Krea 2, LTXVideo2, MageFlow, Z-Image e Wan usam chunks contiguos de *n* blocos nos caminhos whole-block. Outras familias que expoem esta opcao podem continuar usando o comportamento antigo de "checkpoint a cada *n* blocos". Valores maiores podem reduzir recompute, mas normalmente mantem mais activations na VRAM. + +### `--gradient_checkpointing_segment_stride` + +- **O que**: Inicia um segmento com checkpoint a cada *n* blocos nos caminhos segmented whole-block suportados. +- **Exemplo**: Com `--gradient_checkpointing_interval=2` e `--gradient_checkpointing_segment_stride=4`, o SimpleTuner faz checkpoint de dois blocos, executa os dois blocos seguintes normalmente e repete. +- **Nota**: So tem efeito em familias de modelos que expoem suporte segmented whole-block na versao instalada do SimpleTuner. Familias nao suportadas registram um aviso e ignoram o valor. O stride deve ser pelo menos igual ao interval. Veja [Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md). ### `--gradient_checkpointing_backend` @@ -286,7 +299,27 @@ Onde `foo` e seu ambiente de config — ou use `config/config.json` se nao estiv - `torch-ffn`: checkpoint só do lado feed-forward em modelos com fronteira FFN limpa. - `unsloth`: checkpoint do bloco completo compatível e offload dos tensores salvos para CPU. - `unsloth-ffn`: checkpoint só do lado feed-forward e offload dos tensores salvos para CPU. -- **Nota**: So funciona quando `--gradient_checkpointing` esta habilitado. As variantes `unsloth` requerem CUDA. As variantes FFN-only atualmente suportam blocos estilo Flux.1 e MageFlow, e falham explicitamente quando o modelo nao expoe esse escopo. Veja [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) para tradeoffs medidos. +- **Nota**: So funciona quando `--gradient_checkpointing` esta habilitado. As variantes `unsloth` requerem CUDA. As variantes FFN-only atualmente suportam Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan e Z-Image, e falham explicitamente quando o modelo nao expoe esse escopo. Veja [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) para tradeoffs medidos. + +### `--gradient_checkpointing_offload_attention` + +- **O que**: Faz offload para CPU das activations salvas do lado attention em modelos com fronteira attention/FFN limpa. +- **Por que**: Quando transferencia e mais barata que recomputar attention, reduz VRAM sem pagar todo o custo de rematerializar attention. +- **Nota**: Pode ser ativado sozinho. Tambem pode ser combinado com qualquer checkpoint backend que o modelo suportar. So tem efeito em familias de modelos que expoem uma fronteira attention/FFN limpa na versao instalada do SimpleTuner; familias nao suportadas falham explicitamente. + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **Padrao**: `12` +- **O que**: Numero maximo de buckets distintos de tensores CPU pinned usados por activation offload. +- **Por que**: Pinned memory melhora transferencias CPU/GPU, mas resolucoes e comprimentos de texto variaveis podem criar shapes raros. Ao atingir esse limite, novos bucket shapes usam memoria CPU normal. +- **Nota**: Use `0` para desativar o pooling de pinned memory para activation offload. + +### `--gradient_checkpointing_offload_prefetch` + +- **Padrao**: `false` +- **O que**: Aprende a ordem de restore no backward para activations offloaded e prefetch o provavel proximo tensor para a GPU. +- **Por que**: Restore H2D just-in-time quase nao consegue overlap. Com ordem estavel, prefetch pode esconder parte da transferencia atras do backward compute. +- **Nota**: Experimental e ativo apenas com `--gradient_checkpointing_offload_attention`. ### `--refiner_training` @@ -1730,6 +1763,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1744,6 +1778,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2045,6 +2080,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2074,6 +2111,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/OPTIONS.zh.md b/documentation/OPTIONS.zh.md index c9dcfaf33..a72e43fd4 100644 --- a/documentation/OPTIONS.zh.md +++ b/documentation/OPTIONS.zh.md @@ -254,6 +254,13 @@ simpletuner configure config/foo/config.json - **内容**:预训练 Gemma 模型路径或 上的标识符。 - **原因**:训练 Gemma 系模型(例如 LTX-2、Sana、Lumina2)时,可单独指定 Gemma 权重来源,而无需更换基础扩散模型路径。 +### `--qwen_text_encoder_model_name_or_path` + +- **内容**:预训练 Qwen 文本编码器模型路径,或 上的标识符。 +- **默认**:`None`(使用所选模型定义的 Qwen 文本编码器来源)。 +- **原因**:用于在 Qwen 系模型家族中共享或替换 Qwen 文本编码器,而无需编辑 Hugging Face 缓存。 +- **说明**:此选项适用于只有一个 Qwen 文本编码器的模型家族。如果某个模型家族定义了多个 Qwen 编码器,该选项会被忽略,SimpleTuner 会记录警告。 + ### `--max_grounding_entities` - 每张图像用于 GLIGEN 风格空间标注的最大 grounding 实体数。默认值:0(禁用)。典型值:4-16。 @@ -276,8 +283,14 @@ simpletuner configure config/foo/config.json ### `--gradient_checkpointing_interval` -- **内容**:每 *n* 个块进行一次 checkpoint,*n* 必须大于 0。1 等同于启用 `--gradient_checkpointing`,2 则每隔一个块进行一次。 -- **说明**:目前仅 SDXL 与 Flux 支持此选项。SDXL 使用较为权宜的实现。 +- **内容**:transformer block checkpointing 的模型相关 interval。1 基本等同于启用 `--gradient_checkpointing`。 +- **说明**:Flux、Flux.2、Krea 2、LTXVideo2、MageFlow、Z-Image 和 Wan 在 whole-block 路径上使用连续的 *n* 个 block chunk。其他支持此选项的模型族可能仍使用旧的“每第 *n* 个 block checkpoint”行为。值越大可能降低重算开销,但通常会在 VRAM 中保留更多 activation。 + +### `--gradient_checkpointing_segment_stride` + +- **内容**:在支持 segmented whole-block 路径的模型上,每隔 *n* 个 block 启动一个 checkpointed segment。 +- **示例**:设置 `--gradient_checkpointing_interval=2` 且 `--gradient_checkpointing_segment_stride=4` 时,SimpleTuner 会 checkpoint 两个 block,然后正常运行接下来的两个 block,如此重复。 +- **说明**:仅在已安装的 SimpleTuner 版本中暴露 segmented whole-block 支持的模型族上生效。不支持的模型族会记录 warning 并忽略该值。stride 必须大于或等于 interval。参见 [Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md)。 ### `--gradient_checkpointing_backend` @@ -287,7 +300,27 @@ simpletuner configure config/foo/config.json - `torch-ffn`:在有清晰 FFN 边界的模型上,只 checkpoint feed-forward 部分。 - `unsloth`:checkpoint 整个受支持 block,并把保存的 tensor 卸载到 CPU。 - `unsloth-ffn`:只 checkpoint feed-forward 部分,并把保存的 tensor 卸载到 CPU。 -- **说明**:仅在启用 `--gradient_checkpointing` 时生效。`unsloth` 变体需要 CUDA。FFN-only 变体目前支持 Flux.1 风格 blocks 和 MageFlow;模型不支持该 scope 时会直接报错。实测 tradeoff 见 [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md)。 +- **说明**:仅在启用 `--gradient_checkpointing` 时生效。`unsloth` 变体需要 CUDA。FFN-only 变体目前支持 Chroma、Flux、Krea 2、LTXVideo2、MageFlow、Wan 和 Z-Image;模型不支持该 scope 时会直接报错。实测 tradeoff 见 [Unsloth-style checkpointing](experimental/UNSLOTH_CHECKPOINTING.md)。 + +### `--gradient_checkpointing_offload_attention` + +- **内容**:在有清晰 attention/FFN 边界的模型上,把 attention 侧保存的 activations 卸载到 CPU。 +- **原因**:当传输比重算 attention 更便宜时,可降低 VRAM,且不需要支付完整 attention 重算成本。 +- **说明**:它可以单独启用。也可与模型支持的任意 checkpoint backend 组合。仅在已安装的 SimpleTuner 版本中暴露清晰 attention/FFN 边界的模型族上生效;不支持的模型会直接报错。 + +### `--gradient_checkpointing_offload_pin_memory_max_buckets` + +- **默认值**:`12` +- **内容**:activation offload 可使用的 pinned CPU tensor bucket 最大数量。 +- **原因**:pinned memory 有利于 CPU/GPU 传输,但可变分辨率和文本长度会产生少见 tensor 形状。达到上限后,新的 bucket 形状会改用普通 CPU memory。 +- **说明**:设为 `0` 可关闭 activation offload 的 pinned-memory pooling。 + +### `--gradient_checkpointing_offload_prefetch` + +- **默认值**:`false` +- **内容**:学习 offloaded activations 的 backward restore 顺序,并把可能下一个需要的 tensor 预取回 GPU。 +- **原因**:JIT H2D restore 通常很难重叠。顺序稳定后,prefetch 可把部分传输隐藏在 backward compute 后面。 +- **说明**:实验性功能,并且只在启用 `--gradient_checkpointing_offload_attention` 时生效。 ### `--refiner_training` @@ -1737,6 +1770,7 @@ usage: train.py [-h] --model_family [--text_encoder_3_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--text_encoder_4_precision {no_change,int8-quanto,int4-quanto,int2-quanto,int8-torchao,int8dq-torchao,int8dq-int4-torchao,nf4-bnb,int4-torchao,fp8-quanto,fp8uz-quanto,fp8-native,fp8-torchao,fp8wo-torchao,fp8-int4-torchao,fp8-transformerengine}] [--gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL] + [--gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE] [--offload_during_startup [OFFLOAD_DURING_STARTUP]] [--quantize_via {cpu,accelerator,pipeline}] [--quantization_config QUANTIZATION_CONFIG] @@ -1751,6 +1785,7 @@ usage: train.py [-h] --model_family [--pretrained_unet_subfolder PRETRAINED_UNET_SUBFOLDER] [--pretrained_t5_model_name_or_path PRETRAINED_T5_MODEL_NAME_OR_PATH] [--pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH] + [--qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH] [--revision REVISION] [--variant VARIANT] [--base_model_default_dtype {bf16,fp32}] [--unet_attention_slice [UNET_ATTENTION_SLICE]] @@ -2052,6 +2087,8 @@ options: memory. --gradient_checkpointing_interval GRADIENT_CHECKPOINTING_INTERVAL Checkpoint every N transformer blocks + --gradient_checkpointing_segment_stride GRADIENT_CHECKPOINTING_SEGMENT_STRIDE + Start a checkpointed segment every N transformer blocks --offload_during_startup [OFFLOAD_DURING_STARTUP] Offload text encoders to CPU during VAE caching --quantize_via {cpu,accelerator,pipeline} @@ -2081,6 +2118,8 @@ options: Path to pretrained T5 model --pretrained_gemma_model_name_or_path PRETRAINED_GEMMA_MODEL_NAME_OR_PATH Path to pretrained Gemma model + --qwen_text_encoder_model_name_or_path QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH + Path to pretrained Qwen text encoder model --revision REVISION Git branch/tag/commit for model version --variant VARIANT Model variant (e.g., fp16, bf16) --base_model_default_dtype {bf16,fp32} diff --git a/documentation/QUICKSTART.es.md b/documentation/QUICKSTART.es.md index 5083594a6..01870c9e3 100644 --- a/documentation/QUICKSTART.es.md +++ b/documentation/QUICKSTART.es.md @@ -2,55 +2,291 @@ **Nota**: Para configuraciones más avanzadas, consulta el [tutorial](TUTORIAL.md) y la [referencia de opciones](OPTIONS.md). +## Guías de inicio rápido por modelo + +| Modelo | Parámetros | Guía | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.es.md](quickstart/ACE_STEP.es.md) | +| Anima | Not specified | Sin guía dedicada | +| Auraflow | 6B | [AURAFLOW.es.md](quickstart/AURAFLOW.es.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.es.md](quickstart/BOOGU_IMAGE.es.md) | +| Chroma 1 | 8.9B | [CHROMA.es.md](quickstart/CHROMA.es.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.es.md](quickstart/COSMOS2IMAGE.es.md) | +| Cosmos3 | 16B-65B | [COSMOS3.es.md](quickstart/COSMOS3.es.md) | +| DeepFloyd IF | 0.4B-4.3B stages | Sin guía dedicada | +| ERNIE-Image | Not specified | [ERNIE.es.md](quickstart/ERNIE.es.md) | +| Flux.1 | 8B-12B | [FLUX.es.md](quickstart/FLUX.es.md)
[FLUX_KONTEXT.es.md](quickstart/FLUX_KONTEXT.es.md) | +| Flux.2 | 4B-32B | [FLUX2.es.md](quickstart/FLUX2.es.md) | +| HeartMuLa | 3B | [HEARTMULA.es.md](quickstart/HEARTMULA.es.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.es.md](quickstart/HIDREAM.es.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.es.md](quickstart/HUNYUANVIDEO.es.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.es.md](quickstart/IDEOGRAM4.es.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.es.md](quickstart/KANDINSKY5_IMAGE.es.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.es.md](quickstart/KANDINSKY5_VIDEO.es.md) | +| Kwai Kolors | 2.7B | [KOLORS.es.md](quickstart/KOLORS.es.md) | +| Krea2 | Not specified | [KREA2.es.md](quickstart/KREA2.es.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.es.md](quickstart/LONGCAT_IMAGE.es.md)
[LONGCAT_EDIT.es.md](quickstart/LONGCAT_EDIT.es.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.es.md](quickstart/LONGCAT_VIDEO.es.md)
[LONGCAT_VIDEO_EDIT.es.md](quickstart/LONGCAT_VIDEO_EDIT.es.md) | +| LTX Video | ~2.5B | [LTXVIDEO.es.md](quickstart/LTXVIDEO.es.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.es.md](quickstart/LTXVIDEO2.es.md) | +| Lumina2 | 2B | [LUMINA2.es.md](quickstart/LUMINA2.es.md) | +| Mage-Flow | 4B | [MAGEFLOW.es.md](quickstart/MAGEFLOW.es.md) | +| OmniGen | 3.8B | [OMNIGEN.es.md](quickstart/OMNIGEN.es.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.es.md](quickstart/SIGMA.es.md) | +| Qwen Image | 20B | [QWEN_IMAGE.es.md](quickstart/QWEN_IMAGE.es.md)
[QWEN_EDIT.es.md](quickstart/QWEN_EDIT.es.md) | +| Sana | 0.6B-4.8B | [SANA.es.md](quickstart/SANA.es.md) | +| Sana Video | 2B | [SANAVIDEO.es.md](quickstart/SANAVIDEO.es.md) | +| SD 1.x/2.x (Legacy) | 0.9B | Sin guía dedicada | +| Stable Diffusion 3 | 2B-8B | [SD3.es.md](quickstart/SD3.es.md) | +| Stable Diffusion XL | 3.5B | [SDXL.es.md](quickstart/SDXL.es.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.es.md](quickstart/STABLE_CASCADE_C.es.md) | +| Wan Video | 1.3B-14B | [WAN.es.md](quickstart/WAN.es.md) | +| Wan S2V | 14B | [WAN_S2V.es.md](quickstart/WAN_S2V.es.md) | +| Z-Image | 6B | [ZIMAGE.es.md](quickstart/ZIMAGE.es.md) | +| Z-Image Omni | 6B | [ZIMAGE.es.md](quickstart/ZIMAGE.es.md) | +| ZLab I1 | 3B | [ZLAB_i1.es.md](quickstart/ZLAB_i1.es.md) | + ## Compatibilidad de funciones -Para la matriz de funciones completa y más precisa, consulta el [README principal](https://github.com/bghira/SimpleTuner#model-architecture-support). +La matriz completa de compatibilidad se divide por área de función para que cada tabla sea legible. -## Guías de inicio rápido por modelo +
+Soporte de entrenamiento + +| Modelo | PEFT LoRA | LyCORIS | Rango completo | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+Compatibilidad de precisión + +| Modelo | Cuantización | Precisión mixta | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+Granularidad de checkpointing + +| Modelo | Checkpointing de gradiente | Intervalo | Segment stride | Offload de atención | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flujo, destilación y alineación + +| Modelo | Predicción | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+Encoders de texto y tipos de VAE + +| Modelo | Encoders de texto | Parámetros de encoder | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
-| Modelo | Parámetros | PEFT LoRA | Lycoris | Rango completo | Cuantización | Precisión mixta | Checkpointing de gradiente | Flow Shift | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | Guía | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | ✓ | int8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | ✓ | no recomendado | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | ✓ | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✗ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✗ | ✓* | int8 opcional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](quickstart/KREA2.es.md) | -| Boogu-Image 0.1 | - | ✓ | ✓ | ✓* | fp8 opcional | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](quickstart/BOOGU_IMAGE.es.md) | -| zlab i1 | 3B | ✓ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](quickstart/ZLAB_i1.es.md) | -| Ideogram 4 | 9B | ✓ | ✓ | ✓* | fp8 predeterminado, nf4 opcional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](quickstart/IDEOGRAM4.es.md) | -| ACE-Step | 3.5B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | ✓ | int8/fp8 opcional | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | ✓ | no recomendado | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | ✓ | no recomendado | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓ | ✓* | no_change primero | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](quickstart/COSMOS3.es.md) | -| LTX Video | ~2.5B | ✓ | ✓ | ✓ | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓ | ✓* | **requerido** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓ | ✓* | **requerido** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, prior 3.6B | ✓ | ✓ | ✓* | no soportado | fp32 (requerido) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](quickstart/LONGCAT_EDIT.md) | - -*✓ = soportado, ✓* = requiere DeepSpeed/FSDP2 para rango completo, ✗ = no soportado, `✓+` indica que se recomienda checkpointing por presión de VRAM. Ref Inputs marca rutas existentes de condicionamiento por referencia/edición/I2V; `opt` significa opcional y `req` significa requerido por el flavour de edición/I2V. TwinFlow ✓ significa soporte nativo cuando `twinflow_enabled=true` (los modelos de difusión necesitan `diff2flow_enabled+twinflow_allow_diff2flow`). Self-Flow ✓ significa soporte nativo para `crepa_enabled=true` con `crepa_feature_source=self_flow`, `use_ema=true` y `crepa_teacher_block_index` configurado. LayerSync ✓ significa que el backbone expone estados ocultos del transformer para autoalineación; ✗ marca backbones tipo UNet sin ese buffer. †Sliders aplican a LoRA y LyCORIS (incluido LyCORIS de rango completo “full”).* - -> ℹ️ El inicio rápido de Wan incluye presets de las etapas 2.1 y 2.2 y el toggle de time-embedding. Flux Kontext cubre flujos de edición construidos sobre Flux.1. - -> ⚠️ Estos quickstarts son documentos vivos. Espera actualizaciones ocasionales a medida que llegan nuevos modelos o se mejoran recetas de entrenamiento. +*✓ = soportado, ✓* = soportado pero normalmente requiere DeepSpeed/FSDP2 para entrenamiento full-rank, ✗ = no soportado. Ref Inputs marca rutas existentes de condicionamiento por referencia/edición/I2V; `opt` significa opcional y `req` significa requerido por el flavour de edición/I2V.* +*TwinFlow es nativo cuando `twinflow_enabled=true`; los modelos de difusión aún requieren `diff2flow_enabled=true` y `twinflow_allow_diff2flow=true`. Self-Flow se refiere al soporte CREPA self-flow. LayerSync marca backbones que exponen estados ocultos para alineación.* ### Rutas rápidas: Z-Image Turbo y Flux Schnell @@ -76,4 +312,4 @@ Si tu dataset termina con menos muestras utilizables de lo esperado, los archivo - En la WebUI, navega al directorio de tu dataset y selecciónalo para ver estadísticas de filtrado - Revisa los logs durante el procesamiento del dataset para estadísticas como: `Sample processing statistics: {'total_processed': 100, 'skipped': {'too_small': 15, ...}}` -Para solución de problemas detallada, consulta [Solución de problemas de datasets filtrados](DATALOADER.es.md#solución-de-problemas-de-datasets-filtrados) en la documentación del dataloader. +Para solución de problemas detallada, consulta [Solución de problemas de datasets filtrados](DATALOADER.es.md) en la documentación del dataloader. diff --git a/documentation/QUICKSTART.hi.md b/documentation/QUICKSTART.hi.md index 2836acf03..e404d878c 100644 --- a/documentation/QUICKSTART.hi.md +++ b/documentation/QUICKSTART.hi.md @@ -2,55 +2,291 @@ **नोट**: अधिक उन्नत कॉन्फ़िगरेशनों के लिए, [ट्यूटोरियल](TUTORIAL.md) और [options reference](OPTIONS.md) देखें। -## फ़ीचर संगतता +## मॉडल क्विकस्टार्ट गाइड -पूरा और सबसे सटीक फीचर मैट्रिक्स देखने के लिए, [मुख्य README](https://github.com/bghira/SimpleTuner#model-architecture-support) देखें। +| मॉडल | पैरामीटर | गाइड | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.hi.md](quickstart/ACE_STEP.hi.md) | +| Anima | Not specified | Dedicated guide नहीं है | +| Auraflow | 6B | [AURAFLOW.hi.md](quickstart/AURAFLOW.hi.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.hi.md](quickstart/BOOGU_IMAGE.hi.md) | +| Chroma 1 | 8.9B | [CHROMA.hi.md](quickstart/CHROMA.hi.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.hi.md](quickstart/COSMOS2IMAGE.hi.md) | +| Cosmos3 | 16B-65B | [COSMOS3.hi.md](quickstart/COSMOS3.hi.md) | +| DeepFloyd IF | 0.4B-4.3B stages | Dedicated guide नहीं है | +| ERNIE-Image | Not specified | [ERNIE.hi.md](quickstart/ERNIE.hi.md) | +| Flux.1 | 8B-12B | [FLUX.hi.md](quickstart/FLUX.hi.md)
[FLUX_KONTEXT.hi.md](quickstart/FLUX_KONTEXT.hi.md) | +| Flux.2 | 4B-32B | [FLUX2.hi.md](quickstart/FLUX2.hi.md) | +| HeartMuLa | 3B | [HEARTMULA.hi.md](quickstart/HEARTMULA.hi.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.hi.md](quickstart/HIDREAM.hi.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.hi.md](quickstart/HUNYUANVIDEO.hi.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.hi.md](quickstart/IDEOGRAM4.hi.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.hi.md](quickstart/KANDINSKY5_IMAGE.hi.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.hi.md](quickstart/KANDINSKY5_VIDEO.hi.md) | +| Kwai Kolors | 2.7B | [KOLORS.hi.md](quickstart/KOLORS.hi.md) | +| Krea2 | Not specified | [KREA2.hi.md](quickstart/KREA2.hi.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.hi.md](quickstart/LONGCAT_IMAGE.hi.md)
[LONGCAT_EDIT.hi.md](quickstart/LONGCAT_EDIT.hi.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.hi.md](quickstart/LONGCAT_VIDEO.hi.md)
[LONGCAT_VIDEO_EDIT.hi.md](quickstart/LONGCAT_VIDEO_EDIT.hi.md) | +| LTX Video | ~2.5B | [LTXVIDEO.hi.md](quickstart/LTXVIDEO.hi.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.hi.md](quickstart/LTXVIDEO2.hi.md) | +| Lumina2 | 2B | [LUMINA2.hi.md](quickstart/LUMINA2.hi.md) | +| Mage-Flow | 4B | [MAGEFLOW.hi.md](quickstart/MAGEFLOW.hi.md) | +| OmniGen | 3.8B | [OMNIGEN.hi.md](quickstart/OMNIGEN.hi.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.hi.md](quickstart/SIGMA.hi.md) | +| Qwen Image | 20B | [QWEN_IMAGE.hi.md](quickstart/QWEN_IMAGE.hi.md)
[QWEN_EDIT.hi.md](quickstart/QWEN_EDIT.hi.md) | +| Sana | 0.6B-4.8B | [SANA.hi.md](quickstart/SANA.hi.md) | +| Sana Video | 2B | [SANAVIDEO.hi.md](quickstart/SANAVIDEO.hi.md) | +| SD 1.x/2.x (Legacy) | 0.9B | Dedicated guide नहीं है | +| Stable Diffusion 3 | 2B-8B | [SD3.hi.md](quickstart/SD3.hi.md) | +| Stable Diffusion XL | 3.5B | [SDXL.hi.md](quickstart/SDXL.hi.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.hi.md](quickstart/STABLE_CASCADE_C.hi.md) | +| Wan Video | 1.3B-14B | [WAN.hi.md](quickstart/WAN.hi.md) | +| Wan S2V | 14B | [WAN_S2V.hi.md](quickstart/WAN_S2V.hi.md) | +| Z-Image | 6B | [ZIMAGE.hi.md](quickstart/ZIMAGE.hi.md) | +| Z-Image Omni | 6B | [ZIMAGE.hi.md](quickstart/ZIMAGE.hi.md) | +| ZLab I1 | 3B | [ZLAB_i1.hi.md](quickstart/ZLAB_i1.hi.md) | -## मॉडल क्विकस्टार्ट गाइड +## फीचर संगतता + +पूरी compatibility matrix को feature area के अनुसार बांटा गया है ताकि हर table पढ़ने योग्य रहे। + +
+Training support + +| मॉडल | PEFT LoRA | LyCORIS | Full-Rank | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+Precision level support + +| मॉडल | Quantization | Mixed Precision | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+Checkpointing granularity + +| मॉडल | Gradient Checkpoint | Interval | Segment Stride | Attention Offload | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flow, distillation, and alignment + +| मॉडल | Prediction | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+Text encoders and VAE types + +| मॉडल | Text Encoders | Text Encoder Params | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
-| मॉडल | पैरामीटर | PEFT LoRA | Lycoris | फुल-रैंक | क्वांटाइज़ेशन | मिक्स्ड प्रिसिजन | ग्रैड चेकपॉइंट | फ्लो शिफ्ट | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | गाइड | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | ✓ | int8 वैकल्पिक | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | ✓ | int8 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | ✓ | अनुशंसित नहीं | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | ✓ | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✗ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✗ | ✓* | int8 वैकल्पिक | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](quickstart/KREA2.hi.md) | -| Boogu-Image 0.1 | - | ✓ | ✓ | ✓* | fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](quickstart/BOOGU_IMAGE.hi.md) | -| zlab i1 | 3B | ✓ | ✓ | ✓ | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](quickstart/ZLAB_i1.hi.md) | -| Ideogram 4 | 9B | ✓ | ✓ | ✓* | fp8 डिफ़ॉल्ट, nf4 वैकल्पिक | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](quickstart/IDEOGRAM4.hi.md) | -| ACE-Step | 3.5B | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓ | ✓* | int8/fp8/nf4 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | ✓ | int8/fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | ✓ | अनुशंसित नहीं | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | ✓ | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | ✓ | अनुशंसित नहीं | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓ | ✓* | no_change first | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](quickstart/COSMOS3.hi.md) | -| LTX Video | ~2.5B | ✓ | ✓ | ✓ | int8/fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓ | ✓* | int8/fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓ | ✓* | **आवश्यक** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓ | ✓* | **आवश्यक** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, 3.6B prior | ✓ | ✓ | ✓* | समर्थित नहीं | fp32 (आवश्यक) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓ | ✓* | int8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓ | ✓* | int8/fp8 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓ | ✓* | int8/fp8 वैकल्पिक | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓ | ✓* | int8/fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓ | ✓* | int8/fp8 वैकल्पिक | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](quickstart/LONGCAT_EDIT.md) | - -*✓ = समर्थित, ✓* = फुल‑रैंक के लिए DeepSpeed/FSDP2 आवश्यक, ✗ = समर्थित नहीं, `✓+` VRAM दबाव के कारण checkpointing की सिफ़ारिश को दर्शाता है। Ref Inputs सिर्फ मौजूदा reference/edit/I2V conditioning paths को दिखाता है; `opt` वैकल्पिक है और `req` edit/I2V flavour के लिए आवश्यक है। TwinFlow ✓ का अर्थ है `twinflow_enabled=true` होने पर native support (diffusion मॉडल्स को `diff2flow_enabled+twinflow_allow_diff2flow` चाहिए)। Self-Flow ✓ का अर्थ है `crepa_enabled=true`, `crepa_feature_source=self_flow`, `use_ema=true`, और `crepa_teacher_block_index` सेट होने पर native support। LayerSync ✓ का अर्थ है कि backbone self‑alignment के लिए transformer hidden states उपलब्ध कराता है; ✗ UNet‑style backbones को दर्शाता है जिनमें वह buffer नहीं होता। †Sliders LoRA और LyCORIS (full‑rank LyCORIS “full” सहित) पर लागू होते हैं।* - -> ℹ️ Wan quickstart में 2.1 + 2.2 stage presets और time‑embedding toggle शामिल है। Flux Kontext में Flux.1 के ऊपर बने editing वर्कफ़्लो शामिल हैं। - -> ⚠️ ये क्विकस्टार्ट living documents हैं। नए मॉडल आने या प्रशिक्षण रेसिपीज़ सुधरने के साथ समय‑समय पर अपडेट की उम्मीद करें। +*✓ = supported, ✓* = supported लेकिन full-rank training के लिए आम तौर पर DeepSpeed/FSDP2 चाहिए, ✗ = not supported. Ref Inputs मौजूदा reference/edit/I2V conditioning paths को दिखाता है; `opt` optional और `req` edit/I2V flavour में required है.* +*TwinFlow native है जब `twinflow_enabled=true`; diffusion models को अभी भी `diff2flow_enabled=true` और `twinflow_allow_diff2flow=true` चाहिए। Self-Flow CREPA self-flow support है। LayerSync alignment के लिए hidden states expose करने वाले backbones को दिखाता है.* ### तेज़ रास्ते: Z-Image Turbo और Flux Schnell @@ -76,4 +312,4 @@ - WebUI में, अपने dataset directory पर browse करें और filtering statistics देखने के लिए इसे select करें - Dataset processing के दौरान logs में इस तरह के statistics check करें: `Sample processing statistics: {'total_processed': 100, 'skipped': {'too_small': 15, ...}}` -विस्तृत troubleshooting के लिए, dataloader documentation में [Filtered datasets का Troubleshooting](DATALOADER.hi.md#filtered-datasets-का-troubleshooting) देखें। +विस्तृत troubleshooting के लिए, dataloader documentation में [Filtered datasets का Troubleshooting](DATALOADER.hi.md) देखें। diff --git a/documentation/QUICKSTART.ja.md b/documentation/QUICKSTART.ja.md index c3e2fdebe..e26c33129 100644 --- a/documentation/QUICKSTART.ja.md +++ b/documentation/QUICKSTART.ja.md @@ -2,55 +2,291 @@ **注意**: より高度な設定については、[チュートリアル](TUTORIAL.md)および[オプションリファレンス](OPTIONS.md)を参照してください。 +## モデル別クイックスタートガイド + +| モデル | パラメータ | ガイド | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.ja.md](quickstart/ACE_STEP.ja.md) | +| Anima | Not specified | 専用ガイドなし | +| Auraflow | 6B | [AURAFLOW.ja.md](quickstart/AURAFLOW.ja.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.ja.md](quickstart/BOOGU_IMAGE.ja.md) | +| Chroma 1 | 8.9B | [CHROMA.ja.md](quickstart/CHROMA.ja.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.ja.md](quickstart/COSMOS2IMAGE.ja.md) | +| Cosmos3 | 16B-65B | [COSMOS3.ja.md](quickstart/COSMOS3.ja.md) | +| DeepFloyd IF | 0.4B-4.3B stages | 専用ガイドなし | +| ERNIE-Image | Not specified | [ERNIE.ja.md](quickstart/ERNIE.ja.md) | +| Flux.1 | 8B-12B | [FLUX.ja.md](quickstart/FLUX.ja.md)
[FLUX_KONTEXT.ja.md](quickstart/FLUX_KONTEXT.ja.md) | +| Flux.2 | 4B-32B | [FLUX2.ja.md](quickstart/FLUX2.ja.md) | +| HeartMuLa | 3B | [HEARTMULA.ja.md](quickstart/HEARTMULA.ja.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.ja.md](quickstart/HIDREAM.ja.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.ja.md](quickstart/HUNYUANVIDEO.ja.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.ja.md](quickstart/IDEOGRAM4.ja.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.ja.md](quickstart/KANDINSKY5_IMAGE.ja.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.ja.md](quickstart/KANDINSKY5_VIDEO.ja.md) | +| Kwai Kolors | 2.7B | [KOLORS.ja.md](quickstart/KOLORS.ja.md) | +| Krea2 | Not specified | [KREA2.ja.md](quickstart/KREA2.ja.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.ja.md](quickstart/LONGCAT_IMAGE.ja.md)
[LONGCAT_EDIT.ja.md](quickstart/LONGCAT_EDIT.ja.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.ja.md](quickstart/LONGCAT_VIDEO.ja.md)
[LONGCAT_VIDEO_EDIT.ja.md](quickstart/LONGCAT_VIDEO_EDIT.ja.md) | +| LTX Video | ~2.5B | [LTXVIDEO.ja.md](quickstart/LTXVIDEO.ja.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.ja.md](quickstart/LTXVIDEO2.ja.md) | +| Lumina2 | 2B | [LUMINA2.ja.md](quickstart/LUMINA2.ja.md) | +| Mage-Flow | 4B | [MAGEFLOW.ja.md](quickstart/MAGEFLOW.ja.md) | +| OmniGen | 3.8B | [OMNIGEN.ja.md](quickstart/OMNIGEN.ja.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.ja.md](quickstart/SIGMA.ja.md) | +| Qwen Image | 20B | [QWEN_IMAGE.ja.md](quickstart/QWEN_IMAGE.ja.md)
[QWEN_EDIT.ja.md](quickstart/QWEN_EDIT.ja.md) | +| Sana | 0.6B-4.8B | [SANA.ja.md](quickstart/SANA.ja.md) | +| Sana Video | 2B | [SANAVIDEO.ja.md](quickstart/SANAVIDEO.ja.md) | +| SD 1.x/2.x (Legacy) | 0.9B | 専用ガイドなし | +| Stable Diffusion 3 | 2B-8B | [SD3.ja.md](quickstart/SD3.ja.md) | +| Stable Diffusion XL | 3.5B | [SDXL.ja.md](quickstart/SDXL.ja.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.ja.md](quickstart/STABLE_CASCADE_C.ja.md) | +| Wan Video | 1.3B-14B | [WAN.ja.md](quickstart/WAN.ja.md) | +| Wan S2V | 14B | [WAN_S2V.ja.md](quickstart/WAN_S2V.ja.md) | +| Z-Image | 6B | [ZIMAGE.ja.md](quickstart/ZIMAGE.ja.md) | +| Z-Image Omni | 6B | [ZIMAGE.ja.md](quickstart/ZIMAGE.ja.md) | +| ZLab I1 | 3B | [ZLAB_i1.ja.md](quickstart/ZLAB_i1.ja.md) | + ## 機能互換性 -完全かつ最も正確な機能マトリックスについては、[メインREADME](https://github.com/bghira/SimpleTuner#model-architecture-support)を参照してください。 - -## モデルクイックスタートガイド - -| モデル | パラメータ数 | PEFT LoRA | Lycoris | Full-Rank | 量子化 | 混合精度 | Grad Checkpoint | Flow Shift | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | ガイド | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | ✓ | int8 オプション | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | ✓ | int8 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | ✓ | 非推奨 | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | ✓ | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✗ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✗ | ✓* | int8 オプション | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](quickstart/KREA2.ja.md) | -| Boogu-Image 0.1 | - | ✓ | ✓ | ✓* | fp8 オプション | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](quickstart/BOOGU_IMAGE.ja.md) | -| zlab i1 | 3B | ✓ | ✓ | ✓ | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](quickstart/ZLAB_i1.ja.md) | -| Ideogram 4 | 9B | ✓ | ✓ | ✓* | fp8 デフォルト、nf4 オプション | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](quickstart/IDEOGRAM4.ja.md) | -| ACE-Step | 3.5B | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓ | ✓* | int8/fp8/nf4 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | ✓ | int8/fp8 オプション | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | ✓ | 非推奨 | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | ✓ | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | ✓ | 非推奨 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓ | ✓* | no_change first | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](quickstart/COSMOS3.ja.md) | -| LTX Video | ~2.5B | ✓ | ✓ | ✓ | int8/fp8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓ | ✓* | int8/fp8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓ | ✓* | **必須** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓ | ✓* | **必須** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, 3.6B prior | ✓ | ✓ | ✓* | 非対応 | fp32 (必須) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓ | ✓* | int8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓ | ✓* | int8/fp8 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓ | ✓* | int8/fp8 オプション | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓ | ✓* | int8/fp8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓ | ✓* | int8/fp8 オプション | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](quickstart/LONGCAT_EDIT.md) | - -*✓ = サポート、✓* = Full-RankにはDeepSpeed/FSDP2が必要、✗ = 非サポート、`✓+`はVRAMプレッシャーによりチェックポイントが推奨されることを示します。Ref Inputs は既存の参照/編集/I2V 条件パスのみを示し、`opt` は任意、`req` は編集/I2V flavour で必須であることを示します。TwinFlow ✓は`twinflow_enabled=true`のときのネイティブサポートを意味します(拡散モデルには`diff2flow_enabled+twinflow_allow_diff2flow`が必要)。Self-Flow ✓は`crepa_enabled=true`、`crepa_feature_source=self_flow`、`use_ema=true`、および`crepa_teacher_block_index`設定時のネイティブサポートを意味します。LayerSync ✓はバックボーンがセルフアライメント用のトランスフォーマー隠れ状態を公開していることを意味し、✗はそのバッファを持たないUNetスタイルのバックボーンを示します。†SlidersはLoRAおよびLyCORIS(Full-Rank LyCORIS "full"を含む)に適用されます。* - -> ℹ️ Wanクイックスタートには2.1 + 2.2ステージプリセットと時間埋め込みトグルが含まれます。Flux Kontextは、Flux.1をベースに構築された編集ワークフローをカバーします。 - -> ⚠️ これらのクイックスタートは生きたドキュメントです。新しいモデルの登場やトレーニングレシピの改善に伴い、時折更新されることがあります。 +完全な互換性マトリクスは、読みやすいように機能領域ごとに分割しています。 + +
+トレーニングサポート + +| モデル | PEFT LoRA | LyCORIS | フルランク | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+精度レベルのサポート + +| モデル | 量子化 | 混合精度 | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+チェックポイント粒度 + +| モデル | Gradient Checkpoint | Interval | Segment Stride | Attention Offload | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flow・蒸留・アラインメント + +| モデル | Prediction | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+テキストエンコーダーと VAE タイプ + +| モデル | Text Encoders | Text Encoder Params | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
+ +*✓ = サポート、✓* = サポートされていますが、full-rank training では通常 DeepSpeed/FSDP2 が必要、✗ = 未サポート。Ref Inputs は既存の reference/edit/I2V conditioning path を示します。`opt` は任意、`req` は edit/I2V flavour で必須です。* +*TwinFlow は `twinflow_enabled=true` のとき native support です。diffusion models では `diff2flow_enabled=true` と `twinflow_allow_diff2flow=true` も必要です。Self-Flow は CREPA self-flow support です。LayerSync は alignment 用 hidden states を公開する backbone を示します。* ### 高速パス: Z-Image TurboとFlux Schnell @@ -76,4 +312,4 @@ - WebUI でデータセットディレクトリを参照し、選択するとフィルタリング統計が表示されます - データセット処理中のログで次のような統計を確認: `Sample processing statistics: {'total_processed': 100, 'skipped': {'too_small': 15, ...}}` -詳細なトラブルシューティングについては、データローダードキュメントの[フィルタされたデータセットのトラブルシューティング](DATALOADER.ja.md#フィルタされたデータセットのトラブルシューティング)を参照してください。 +詳細なトラブルシューティングについては、データローダードキュメントの[フィルタされたデータセットのトラブルシューティング](DATALOADER.ja.md)を参照してください。 diff --git a/documentation/QUICKSTART.md b/documentation/QUICKSTART.md index 63cdbe2a7..240dacf64 100644 --- a/documentation/QUICKSTART.md +++ b/documentation/QUICKSTART.md @@ -2,55 +2,291 @@ **Note**: For more advanced configurations, see the [tutorial](TUTORIAL.md) and [options reference](OPTIONS.md). +## Model Quickstart Guides + +| Model | Parameters | Guide | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.md](/documentation/quickstart/ACE_STEP.md) | +| Anima | Not specified | No dedicated guide | +| Auraflow | 6B | [AURAFLOW.md](/documentation/quickstart/AURAFLOW.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.md](/documentation/quickstart/BOOGU_IMAGE.md) | +| Chroma 1 | 8.9B | [CHROMA.md](/documentation/quickstart/CHROMA.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.md](/documentation/quickstart/COSMOS2IMAGE.md) | +| Cosmos3 | 16B-65B | [COSMOS3.md](/documentation/quickstart/COSMOS3.md) | +| DeepFloyd IF | 0.4B-4.3B stages | No dedicated guide | +| ERNIE-Image | Not specified | [ERNIE.md](/documentation/quickstart/ERNIE.md) | +| Flux.1 | 8B-12B | [FLUX.md](/documentation/quickstart/FLUX.md)
[FLUX_KONTEXT.md](/documentation/quickstart/FLUX_KONTEXT.md) | +| Flux.2 | 4B-32B | [FLUX2.md](/documentation/quickstart/FLUX2.md) | +| HeartMuLa | 3B | [HEARTMULA.md](/documentation/quickstart/HEARTMULA.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.md](/documentation/quickstart/HIDREAM.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.md](/documentation/quickstart/HUNYUANVIDEO.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.md](/documentation/quickstart/IDEOGRAM4.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.md](/documentation/quickstart/KANDINSKY5_IMAGE.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.md](/documentation/quickstart/KANDINSKY5_VIDEO.md) | +| Kwai Kolors | 2.7B | [KOLORS.md](/documentation/quickstart/KOLORS.md) | +| Krea2 | Not specified | [KREA2.md](/documentation/quickstart/KREA2.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.md](/documentation/quickstart/LONGCAT_IMAGE.md)
[LONGCAT_EDIT.md](/documentation/quickstart/LONGCAT_EDIT.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.md](/documentation/quickstart/LONGCAT_VIDEO.md)
[LONGCAT_VIDEO_EDIT.md](/documentation/quickstart/LONGCAT_VIDEO_EDIT.md) | +| LTX Video | ~2.5B | [LTXVIDEO.md](/documentation/quickstart/LTXVIDEO.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.md](/documentation/quickstart/LTXVIDEO2.md) | +| Lumina2 | 2B | [LUMINA2.md](/documentation/quickstart/LUMINA2.md) | +| Mage-Flow | 4B | [MAGEFLOW.md](/documentation/quickstart/MAGEFLOW.md) | +| OmniGen | 3.8B | [OMNIGEN.md](/documentation/quickstart/OMNIGEN.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.md](/documentation/quickstart/SIGMA.md) | +| Qwen Image | 20B | [QWEN_IMAGE.md](/documentation/quickstart/QWEN_IMAGE.md)
[QWEN_EDIT.md](/documentation/quickstart/QWEN_EDIT.md) | +| Sana | 0.6B-4.8B | [SANA.md](/documentation/quickstart/SANA.md) | +| Sana Video | 2B | [SANAVIDEO.md](/documentation/quickstart/SANAVIDEO.md) | +| SD 1.x/2.x (Legacy) | 0.9B | No dedicated guide | +| Stable Diffusion 3 | 2B-8B | [SD3.md](/documentation/quickstart/SD3.md) | +| Stable Diffusion XL | 3.5B | [SDXL.md](/documentation/quickstart/SDXL.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.md](/documentation/quickstart/STABLE_CASCADE_C.md) | +| Wan Video | 1.3B-14B | [WAN.md](/documentation/quickstart/WAN.md) | +| Wan S2V | 14B | [WAN_S2V.md](/documentation/quickstart/WAN_S2V.md) | +| Z-Image | 6B | [ZIMAGE.md](/documentation/quickstart/ZIMAGE.md) | +| Z-Image Omni | 6B | [ZIMAGE.md](/documentation/quickstart/ZIMAGE.md) | +| ZLab I1 | 3B | [ZLAB_i1.md](/documentation/quickstart/ZLAB_i1.md) | + ## Feature Compatibility -For the complete and most accurate feature matrix, refer to the [main README](https://github.com/bghira/SimpleTuner#model-architecture-support). +The complete compatibility matrix is split by feature area so each table stays readable. -## Model Quickstart Guides +
+Training support + +| Model | PEFT LoRA | LyCORIS | Full-Rank | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+Precision level support + +| Model | Quantization | Mixed Precision | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+Checkpointing granularity + +| Model | Gradient Checkpoint | Interval | Segment Stride | Attention Offload | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flow, distillation, and alignment + +| Model | Prediction | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+Text encoders and VAE types + +| Model | Text Encoders | Text Encoder Params | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
-| Model | Params | PEFT LoRA | Full-Rank | Quantization | Mixed Precision | Grad Checkpoint | Flow Shift | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | Guide | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | int8 optional | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](/documentation/quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | int8 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](/documentation/quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | not recommended | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](/documentation/quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](/documentation/quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](/documentation/quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](/documentation/quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](/documentation/quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](/documentation/quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✓* | int8 optional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](/documentation/quickstart/KREA2.md) | -| Boogu-Image 0.1 | - | ✓ | ✓* | fp8 optional | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](/documentation/quickstart/BOOGU_IMAGE.md) | -| zlab i1 | 3B | ✓ | ✓ | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](/documentation/quickstart/ZLAB_i1.md) | -| Ideogram 4 | 9B | ✓ | ✓* | fp8 default, nf4 optional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](/documentation/quickstart/IDEOGRAM4.md) | -| ACE-Step | 3.5B | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](/documentation/quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](/documentation/quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](/documentation/quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓* | int8/fp8/nf4 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](/documentation/quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | int8/fp8 optional | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](/documentation/quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | not recommended | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](/documentation/quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](/documentation/quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | not recommended | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](/documentation/quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓* | no_change first | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](/documentation/quickstart/COSMOS3.md) | -| LTX Video | ~2.5B | ✓ | ✓ | int8/fp8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](/documentation/quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓* | int8/fp8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](/documentation/quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](/documentation/quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](/documentation/quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](/documentation/quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓* | **required** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](/documentation/quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓* | **required** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](/documentation/quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, 3.6B prior | ✓ | ✓* | not supported | fp32 (required) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](/documentation/quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](/documentation/quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓* | int8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](/documentation/quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓* | int8/fp8 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](/documentation/quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓* | int8/fp8 optional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](/documentation/quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓* | int8/fp8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](/documentation/quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓* | int8/fp8 optional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](/documentation/quickstart/LONGCAT_EDIT.md) | - -*✓ = supported, ✓* = requires DeepSpeed/FSDP2 for full-rank, ✗ = not supported, `✓+` indicates checkpointing is recommended due to VRAM pressure. Ref Inputs marks existing reference/edit/I2V conditioning paths; `opt` means optional, `req` means the edit/I2V flavour requires it. TwinFlow ✓ means native support when `twinflow_enabled=true` (diffusion models need `diff2flow_enabled+twinflow_allow_diff2flow`). Self-Flow ✓ means native support for `crepa_enabled=true` with `crepa_feature_source=self_flow`, `use_ema=true`, and `crepa_teacher_block_index` set. LayerSync ✓ means the backbone exposes transformer hidden states for self-alignment; ✗ marks UNet-style backbones without that buffer. †Sliders apply to LoRA and LyCORIS (including full-rank LyCORIS "full"). All models support LyCORIS.* - -> ℹ️ Wan quickstart includes 2.1 + 2.2 stage presets and the time-embedding toggle. Flux Kontext covers editing workflows built atop Flux.1. - -> ⚠️ These quickstarts are living documents. Expect occasional updates as new models land or training recipes improve. +*✓ = supported, ✓* = supported but usually requires DeepSpeed/FSDP2 for full-rank training, ✗ = not supported. Ref Inputs marks existing reference/edit/I2V conditioning paths; `opt` means optional and `req` means the edit/I2V flavour requires it.* +*TwinFlow support is native when `twinflow_enabled=true`; diffusion models still require `diff2flow_enabled=true` plus `twinflow_allow_diff2flow=true`. Self-Flow refers to CREPA self-flow support. LayerSync marks backbones that expose hidden states for alignment.* ### Fast paths: Z-Image Turbo & Flux Schnell diff --git a/documentation/QUICKSTART.pt-BR.md b/documentation/QUICKSTART.pt-BR.md index 753962ece..370599c88 100644 --- a/documentation/QUICKSTART.pt-BR.md +++ b/documentation/QUICKSTART.pt-BR.md @@ -2,55 +2,291 @@ **Nota**: Para configurações mais avançadas, veja o [tutorial](TUTORIAL.md) e a [referência de opções](OPTIONS.md). +## Guias de início rápido por modelo + +| Modelo | Parâmetros | Guia | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.pt-BR.md](quickstart/ACE_STEP.pt-BR.md) | +| Anima | Not specified | Sem guia dedicada | +| Auraflow | 6B | [AURAFLOW.pt-BR.md](quickstart/AURAFLOW.pt-BR.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.pt-BR.md](quickstart/BOOGU_IMAGE.pt-BR.md) | +| Chroma 1 | 8.9B | [CHROMA.pt-BR.md](quickstart/CHROMA.pt-BR.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.pt-BR.md](quickstart/COSMOS2IMAGE.pt-BR.md) | +| Cosmos3 | 16B-65B | [COSMOS3.pt-BR.md](quickstart/COSMOS3.pt-BR.md) | +| DeepFloyd IF | 0.4B-4.3B stages | Sem guia dedicada | +| ERNIE-Image | Not specified | [ERNIE.pt-BR.md](quickstart/ERNIE.pt-BR.md) | +| Flux.1 | 8B-12B | [FLUX.pt-BR.md](quickstart/FLUX.pt-BR.md)
[FLUX_KONTEXT.pt-BR.md](quickstart/FLUX_KONTEXT.pt-BR.md) | +| Flux.2 | 4B-32B | [FLUX2.pt-BR.md](quickstart/FLUX2.pt-BR.md) | +| HeartMuLa | 3B | [HEARTMULA.pt-BR.md](quickstart/HEARTMULA.pt-BR.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.pt-BR.md](quickstart/HIDREAM.pt-BR.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.pt-BR.md](quickstart/HUNYUANVIDEO.pt-BR.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.pt-BR.md](quickstart/IDEOGRAM4.pt-BR.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.pt-BR.md](quickstart/KANDINSKY5_IMAGE.pt-BR.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.pt-BR.md](quickstart/KANDINSKY5_VIDEO.pt-BR.md) | +| Kwai Kolors | 2.7B | [KOLORS.pt-BR.md](quickstart/KOLORS.pt-BR.md) | +| Krea2 | Not specified | [KREA2.pt-BR.md](quickstart/KREA2.pt-BR.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.pt-BR.md](quickstart/LONGCAT_IMAGE.pt-BR.md)
[LONGCAT_EDIT.pt-BR.md](quickstart/LONGCAT_EDIT.pt-BR.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.pt-BR.md](quickstart/LONGCAT_VIDEO.pt-BR.md)
[LONGCAT_VIDEO_EDIT.pt-BR.md](quickstart/LONGCAT_VIDEO_EDIT.pt-BR.md) | +| LTX Video | ~2.5B | [LTXVIDEO.pt-BR.md](quickstart/LTXVIDEO.pt-BR.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.pt-BR.md](quickstart/LTXVIDEO2.pt-BR.md) | +| Lumina2 | 2B | [LUMINA2.pt-BR.md](quickstart/LUMINA2.pt-BR.md) | +| Mage-Flow | 4B | [MAGEFLOW.pt-BR.md](quickstart/MAGEFLOW.pt-BR.md) | +| OmniGen | 3.8B | [OMNIGEN.pt-BR.md](quickstart/OMNIGEN.pt-BR.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.pt-BR.md](quickstart/SIGMA.pt-BR.md) | +| Qwen Image | 20B | [QWEN_IMAGE.pt-BR.md](quickstart/QWEN_IMAGE.pt-BR.md)
[QWEN_EDIT.pt-BR.md](quickstart/QWEN_EDIT.pt-BR.md) | +| Sana | 0.6B-4.8B | [SANA.pt-BR.md](quickstart/SANA.pt-BR.md) | +| Sana Video | 2B | [SANAVIDEO.pt-BR.md](quickstart/SANAVIDEO.pt-BR.md) | +| SD 1.x/2.x (Legacy) | 0.9B | Sem guia dedicada | +| Stable Diffusion 3 | 2B-8B | [SD3.pt-BR.md](quickstart/SD3.pt-BR.md) | +| Stable Diffusion XL | 3.5B | [SDXL.pt-BR.md](quickstart/SDXL.pt-BR.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.pt-BR.md](quickstart/STABLE_CASCADE_C.pt-BR.md) | +| Wan Video | 1.3B-14B | [WAN.pt-BR.md](quickstart/WAN.pt-BR.md) | +| Wan S2V | 14B | [WAN_S2V.pt-BR.md](quickstart/WAN_S2V.pt-BR.md) | +| Z-Image | 6B | [ZIMAGE.pt-BR.md](quickstart/ZIMAGE.pt-BR.md) | +| Z-Image Omni | 6B | [ZIMAGE.pt-BR.md](quickstart/ZIMAGE.pt-BR.md) | +| ZLab I1 | 3B | [ZLAB_i1.pt-BR.md](quickstart/ZLAB_i1.pt-BR.md) | + ## Compatibilidade de recursos -Para a matriz completa e mais precisa de recursos, consulte o [README principal](https://github.com/bghira/SimpleTuner#model-architecture-support). +A matriz completa de compatibilidade é dividida por área de recurso para manter cada tabela legível. -## Guias de início rápido por modelo +
+Suporte de treinamento + +| Modelo | PEFT LoRA | LyCORIS | Full-Rank | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+Suporte a níveis de precisão + +| Modelo | Quantização | Precisão mista | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+Granularidade de checkpointing + +| Modelo | Gradient Checkpoint | Interval | Segment Stride | Attention Offload | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flow, destilação e alinhamento + +| Modelo | Prediction | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+Text encoders e tipos de VAE + +| Modelo | Text Encoders | Text Encoder Params | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
-| Modelo | Parâmetros | LoRA PEFT | Lycoris | Full-Rank | Quantização | Precisão mista | Checkpoint de gradiente | Flow Shift | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | Guia | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | ✓ | int8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | ✓ | não recomendado | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | ✓ | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✗ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✗ | ✓* | int8 opcional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](quickstart/KREA2.pt-BR.md) | -| Boogu-Image 0.1 | - | ✓ | ✓ | ✓* | fp8 opcional | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](quickstart/BOOGU_IMAGE.pt-BR.md) | -| zlab i1 | 3B | ✓ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](quickstart/ZLAB_i1.pt-BR.md) | -| Ideogram 4 | 9B | ✓ | ✓ | ✓* | fp8 padrão, nf4 opcional | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](quickstart/IDEOGRAM4.pt-BR.md) | -| ACE-Step | 3.5B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓ | ✓* | int8/fp8/nf4 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | ✓ | int8/fp8 opcional | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | ✓ | não recomendado | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | ✓ | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | ✓ | não recomendado | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓ | ✓* | no_change primeiro | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](quickstart/COSMOS3.pt-BR.md) | -| LTX Video | ~2.5B | ✓ | ✓ | ✓ | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓ | ✓* | **obrigatório** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓ | ✓* | **obrigatório** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, 3.6B prior | ✓ | ✓ | ✓* | não suportado | fp32 (obrigatório) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓ | ✓* | int8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓ | ✓* | int8/fp8 opcional | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](quickstart/LONGCAT_EDIT.md) | - -*✓ = suportado, ✓* = requer DeepSpeed/FSDP2 para full-rank, ✗ = não suportado, `✓+` indica que o checkpointing é recomendado devido à pressão de VRAM. Ref Inputs marca caminhos existentes de condicionamento por referência/edição/I2V; `opt` significa opcional e `req` significa obrigatório para o flavour de edição/I2V. TwinFlow ✓ significa suporte nativo quando `twinflow_enabled=true` (modelos de difusão precisam de `diff2flow_enabled+twinflow_allow_diff2flow`). Self-Flow ✓ significa suporte nativo para `crepa_enabled=true` com `crepa_feature_source=self_flow`, `use_ema=true` e `crepa_teacher_block_index` definido. LayerSync ✓ significa que o backbone expõe estados ocultos do transformer para autoalinhamento; ✗ marca backbones estilo UNet sem esse buffer. †Sliders se aplicam a LoRA e LyCORIS (incluindo LyCORIS full-rank “full”).* - -> ℹ️ O quickstart do Wan inclui presets das etapas 2.1 + 2.2 e o toggle de time-embedding. Flux Kontext cobre fluxos de edição construídos sobre o Flux.1. - -> ⚠️ Estes quickstarts são documentos vivos. Espere atualizações ocasionais conforme novos modelos chegam ou as receitas de treinamento melhoram. +*✓ = suportado, ✓* = suportado mas normalmente requer DeepSpeed/FSDP2 para treinamento full-rank, ✗ = não suportado. Ref Inputs marca rotas existentes de condicionamento por referência/edição/I2V; `opt` significa opcional e `req` significa obrigatório para o flavour de edição/I2V.* +*TwinFlow é nativo quando `twinflow_enabled=true`; modelos de difusão ainda exigem `diff2flow_enabled=true` e `twinflow_allow_diff2flow=true`. Self-Flow se refere ao suporte CREPA self-flow. LayerSync marca backbones que expõem hidden states para alinhamento.* ### Caminhos rápidos: Z-Image Turbo e Flux Schnell @@ -76,4 +312,4 @@ Se seu dataset acaba com menos amostras utilizáveis do que você esperava, arqu - Na WebUI, navegue até o diretório do seu dataset e selecione-o para ver estatísticas de filtragem - Verifique os logs durante o processamento do dataset por estatísticas como: `Sample processing statistics: {'total_processed': 100, 'skipped': {'too_small': 15, ...}}` -Para solução de problemas detalhada, consulte [Solucionando problemas de datasets filtrados](DATALOADER.pt-BR.md#solucionando-problemas-de-datasets-filtrados) na documentação do dataloader. +Para solução de problemas detalhada, consulte [Solucionando problemas de datasets filtrados](DATALOADER.pt-BR.md) na documentação do dataloader. diff --git a/documentation/QUICKSTART.zh.md b/documentation/QUICKSTART.zh.md index 6b12ad619..47c3366f7 100644 --- a/documentation/QUICKSTART.zh.md +++ b/documentation/QUICKSTART.zh.md @@ -2,55 +2,291 @@ **注意**:如需更高级的配置,请参阅[教程](TUTORIAL.md)和[选项参考](OPTIONS.md)。 +## 模型快速开始指南 + +| 模型 | 参数量 | 指南 | +| --- | --- | --- | +| ACE-Step | 3.5B | [ACE_STEP.zh.md](quickstart/ACE_STEP.zh.md) | +| Anima | Not specified | 暂无专用指南 | +| Auraflow | 6B | [AURAFLOW.zh.md](quickstart/AURAFLOW.zh.md) | +| Boogu-Image | Not specified | [BOOGU_IMAGE.zh.md](quickstart/BOOGU_IMAGE.zh.md) | +| Chroma 1 | 8.9B | [CHROMA.zh.md](quickstart/CHROMA.zh.md) | +| Cosmos2 | 2B-14B | [COSMOS2IMAGE.zh.md](quickstart/COSMOS2IMAGE.zh.md) | +| Cosmos3 | 16B-65B | [COSMOS3.zh.md](quickstart/COSMOS3.zh.md) | +| DeepFloyd IF | 0.4B-4.3B stages | 暂无专用指南 | +| ERNIE-Image | Not specified | [ERNIE.zh.md](quickstart/ERNIE.zh.md) | +| Flux.1 | 8B-12B | [FLUX.zh.md](quickstart/FLUX.zh.md)
[FLUX_KONTEXT.zh.md](quickstart/FLUX_KONTEXT.zh.md) | +| Flux.2 | 4B-32B | [FLUX2.zh.md](quickstart/FLUX2.zh.md) | +| HeartMuLa | 3B | [HEARTMULA.zh.md](quickstart/HEARTMULA.zh.md) | +| HiDream | 17B (8.5B MoE) | [HIDREAM.zh.md](quickstart/HIDREAM.zh.md) | +| Hunyuan Video | 8.3B | [HUNYUANVIDEO.zh.md](quickstart/HUNYUANVIDEO.zh.md) | +| Ideogram 4 | 9B | [IDEOGRAM4.zh.md](quickstart/IDEOGRAM4.zh.md) | +| Kandinsky 5.0 Image | 6B (lite) | [KANDINSKY5_IMAGE.zh.md](quickstart/KANDINSKY5_IMAGE.zh.md) | +| Kandinsky 5.0 Video | 2B lite, 19B pro | [KANDINSKY5_VIDEO.zh.md](quickstart/KANDINSKY5_VIDEO.zh.md) | +| Kwai Kolors | 2.7B | [KOLORS.zh.md](quickstart/KOLORS.zh.md) | +| Krea2 | Not specified | [KREA2.zh.md](quickstart/KREA2.zh.md) | +| LongCat Image | 6B | [LONGCAT_IMAGE.zh.md](quickstart/LONGCAT_IMAGE.zh.md)
[LONGCAT_EDIT.zh.md](quickstart/LONGCAT_EDIT.zh.md) | +| LongCat Video | 13.6B | [LONGCAT_VIDEO.zh.md](quickstart/LONGCAT_VIDEO.zh.md)
[LONGCAT_VIDEO_EDIT.zh.md](quickstart/LONGCAT_VIDEO_EDIT.zh.md) | +| LTX Video | ~2.5B | [LTXVIDEO.zh.md](quickstart/LTXVIDEO.zh.md) | +| LTX Video 2 | 19B | [LTXVIDEO2.zh.md](quickstart/LTXVIDEO2.zh.md) | +| Lumina2 | 2B | [LUMINA2.zh.md](quickstart/LUMINA2.zh.md) | +| Mage-Flow | 4B | [MAGEFLOW.zh.md](quickstart/MAGEFLOW.zh.md) | +| OmniGen | 3.8B | [OMNIGEN.zh.md](quickstart/OMNIGEN.zh.md) | +| PixArt Sigma | 0.6B-0.9B | [SIGMA.zh.md](quickstart/SIGMA.zh.md) | +| Qwen Image | 20B | [QWEN_IMAGE.zh.md](quickstart/QWEN_IMAGE.zh.md)
[QWEN_EDIT.zh.md](quickstart/QWEN_EDIT.zh.md) | +| Sana | 0.6B-4.8B | [SANA.zh.md](quickstart/SANA.zh.md) | +| Sana Video | 2B | [SANAVIDEO.zh.md](quickstart/SANAVIDEO.zh.md) | +| SD 1.x/2.x (Legacy) | 0.9B | 暂无专用指南 | +| Stable Diffusion 3 | 2B-8B | [SD3.zh.md](quickstart/SD3.zh.md) | +| Stable Diffusion XL | 3.5B | [SDXL.zh.md](quickstart/SDXL.zh.md) | +| Stable Cascade (Stage C) | 1B, 3.6B prior | [STABLE_CASCADE_C.zh.md](quickstart/STABLE_CASCADE_C.zh.md) | +| Wan Video | 1.3B-14B | [WAN.zh.md](quickstart/WAN.zh.md) | +| Wan S2V | 14B | [WAN_S2V.zh.md](quickstart/WAN_S2V.zh.md) | +| Z-Image | 6B | [ZIMAGE.zh.md](quickstart/ZIMAGE.zh.md) | +| Z-Image Omni | 6B | [ZIMAGE.zh.md](quickstart/ZIMAGE.zh.md) | +| ZLab I1 | 3B | [ZLAB_i1.zh.md](quickstart/ZLAB_i1.zh.md) | + ## 功能兼容性 -完整且最准确的功能矩阵,请参阅[主 README](https://github.com/bghira/SimpleTuner#model-architecture-support)。 +完整兼容性矩阵按功能领域拆分,以保持每张表易读。 -## 模型快速开始指南 +
+训练支持 + +| 模型 | PEFT LoRA | LyCORIS | 全秩 | ControlNet | Ref Inputs | +| --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓* | ✗ | ✗ | +| Anima | ✓ | ✓ | ✓* | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓* | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| Chroma 1 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓* | ✗ | audio opt | +| DeepFloyd IF | ✓ | ✓ | ✓ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Flux.1 | ✓ | ✓ | ✓* | ✓ | ✓ opt (Kontext) | +| Flux.2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| HeartMuLa | ✓ | ✓ | ✓* | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓* | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Ideogram 4 | ✓ | ✓ | ✓* | ✗ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓* | ✗ | ✓ I2I | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V | +| Kwai Kolors | ✓ | ✓ | ✓ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| LongCat Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| LongCat Video | ✓ | ✓ | ✓* | ✗ | ✓ opt/edit | +| LTX Video | ✓ | ✓ | ✓ | ✗ | ✓ I2V | +| LTX Video 2 | ✓ | ✓ | ✓* | ✗ | ✓ opt | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓* | ✗ | ✓ edit | +| OmniGen | ✓ | ✓ | ✓ | ✗ | ✗ | +| PixArt Sigma | ✗ | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓* | ✗ | ✓ req (Edit) | +| Sana | ✗ | ✓ | ✓ | ✗ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓* | ✓ | ✗ | +| Stable Diffusion XL | ✓ | ✓ | ✓ | ✓ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓* | ✗ | ✗ | +| Wan Video | ✓ | ✓ | ✓* | ✗ | ✓ I2V/VACE | +| Wan S2V | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image | ✓ | ✓ | ✓* | ✗ | ✗ | +| Z-Image Omni | ✓ | ✓ | ✓* | ✗ | ✓ opt (Edit) | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | ✗ | + +
+ +
+精度级别支持 + +| 模型 | 量化 | 混合精度 | +| --- | --- | --- | +| ACE-Step | int8 optional | bf16 | +| Anima | not specified | bf16 | +| Auraflow | int8/fp8/nf4 optional | bf16 | +| Boogu-Image | fp8 optional | bf16 | +| Chroma 1 | int8/fp8/nf4 optional | bf16 | +| Cosmos2 | int8 optional | bf16 | +| Cosmos3 | no_change first; int8 optional | bf16 | +| DeepFloyd IF | not recommended | bf16 | +| ERNIE-Image | int8 optional | bf16 | +| Flux.1 | int8/fp8/nf4 optional | bf16 | +| Flux.2 | int8/fp8/nf4 optional | bf16 | +| HeartMuLa | int8 optional | bf16 | +| HiDream | int8/fp8/nf4 optional | bf16 | +| Hunyuan Video | int8 optional | bf16 | +| Ideogram 4 | fp8 default, nf4 optional | bf16 | +| Kandinsky 5.0 Image | int8 optional | bf16 | +| Kandinsky 5.0 Video | int8 optional | bf16 | +| Kwai Kolors | not recommended | bf16 | +| Krea2 | int8 optional | bf16 | +| LongCat Image | int8/fp8 optional | bf16 | +| LongCat Video | int8/fp8 optional | bf16 | +| LTX Video | int8/fp8 optional | bf16 | +| LTX Video 2 | int8/fp8 optional | bf16 | +| Lumina2 | int8 optional | bf16 | +| Mage-Flow | fp8 optional | bf16 | +| OmniGen | int8/fp8 optional | bf16 | +| PixArt Sigma | int8 optional | bf16 | +| Qwen Image | required (int8/nf4) | bf16 | +| Sana | int8 optional | bf16 | +| Sana Video | not recommended for full | bf16 | +| SD 1.x/2.x (Legacy) | int8/nf4 optional | bf16 | +| Stable Diffusion 3 | int8/fp8/nf4 optional | bf16 | +| Stable Diffusion XL | int8/nf4 optional | bf16 | +| Stable Cascade (Stage C) | not supported | fp32 required | +| Wan Video | int8 optional | bf16 | +| Wan S2V | int8 optional | bf16 | +| Z-Image | int8 optional | bf16 | +| Z-Image Omni | int8 optional | bf16 | +| ZLab I1 | int8 optional | bf16 | + +
+ +
+检查点粒度 + +| 模型 | Gradient Checkpoint | Interval | Segment Stride | Attention Offload | +| --- | :---: | :---: | :---: | :---: | +| ACE-Step | ✓ | ✓ | ✓ | ✗ | +| Anima | ✓ | ✗ | ✗ | ✗ | +| Auraflow | ✓ | ✓ | ✓ | ✗ | +| Boogu-Image | ✓ | ✓ | ✓ | ✗ | +| Chroma 1 | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | ✓ | ✓ | ✓ | ✗ | +| Cosmos3 | ✓ | ✓ | ✓ | ✗ | +| DeepFloyd IF | ✓ | ✗ | ✗ | ✗ | +| ERNIE-Image | ✓ | ✓ | ✓ | ✗ | +| Flux.1 | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | ✓ | ✗ | ✗ | ✗ | +| HiDream | ✓ | ✓ | ✓ | ✗ | +| Hunyuan Video | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | ✓ | ✓ | ✓ | ✗ | +| Kandinsky 5.0 Image | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | ✓ | ✗ | ✗ | ✗ | +| Krea2 | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | ✓ | ✓ | ✓ | ✓ | +| LTX Video | ✓ | ✓ | ✓ | ✗ | +| LTX Video 2 | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | ✓ | ✓ | ✓ | ✗ | +| Mage-Flow | ✓ | ✓ | ✓ | ✓ | +| OmniGen | ✓ | ✗ | ✗ | ✗ | +| PixArt Sigma | ✓ | ✓ | ✓ | ✗ | +| Qwen Image | ✓ | ✓ | ✓ | ✗ | +| Sana | ✓ | ✓ | ✓ | ✗ | +| Sana Video | ✓ | ✓ | ✓ | ✗ | +| SD 1.x/2.x (Legacy) | ✓ | ✗ | ✗ | ✗ | +| Stable Diffusion 3 | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | ✓ | ✗ | ✗ | ✗ | +| Stable Cascade (Stage C) | ✓ | ✓ | ✓ | ✗ | +| Wan Video | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | ✓ | ✓ | ✓ | ✗ | +| Z-Image | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | ✓ | ✗ | ✗ | ✗ | +| ZLab I1 | ✓ | ✓ | ✓ | ✗ | + +
+ +
+Flow、蒸馏与对齐 + +| 模型 | Prediction | Flow Shift | TwinFlow | Self-Flow | LayerSync | Sliders | +| --- | --- | :---: | :---: | :---: | :---: | :---: | +| ACE-Step | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Anima | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Auraflow | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Boogu-Image | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Chroma 1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Cosmos2 | sample | ✗ | ✗ | ✓ | ✓ | ✓ | +| Cosmos3 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| DeepFloyd IF | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| ERNIE-Image | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| Flux.1 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Flux.2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| HeartMuLa | autoregressive next-token | ✗ | ✗ | ✗ | ✗ | ✗ | +| HiDream | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Hunyuan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Ideogram 4 | flow matching | ✓ | ✗ | ✗ | ✗ | ✓ | +| Kandinsky 5.0 Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kandinsky 5.0 Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Kwai Kolors | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Krea2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LongCat Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| LTX Video 2 | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Lumina2 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Mage-Flow | flow matching | ✓ | ✓ | ✗ | ✓ | ✓ | +| OmniGen | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| PixArt Sigma | epsilon | ✗ | ✗ | ✓ | ✓ | ✓ | +| Qwen Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Sana Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| SD 1.x/2.x (Legacy) | epsilon / v-pred | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Diffusion 3 | flow matching | ✓ (SLG) | ✓ | ✓ | ✓ | ✓ | +| Stable Diffusion XL | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Stable Cascade (Stage C) | epsilon | ✗ | ✗ | ✗ | ✗ | ✓ | +| Wan Video | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Wan S2V | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | +| Z-Image | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| Z-Image Omni | flow matching | ✓ | ✓ | ✓ | ✓ | ✓ | +| ZLab I1 | flow matching | ✓ | ✗ | ✓ | ✓ | ✓ | + +
+ +
+文本编码器与 VAE 类型 + +| 模型 | Text Encoders | Text Encoder Params | VAE | +| --- | --- | --- | --- | +| ACE-Step | UMT5 Encoder | 0.6B | Music DCAE | +| Anima | Qwen3 0.6B | 0.6B | Qwen Image VAE | +| Auraflow | Pile T5 | not specified | AutoencoderKL | +| Boogu-Image | Qwen3-VL | not specified | AutoencoderKL | +| Chroma 1 | T5 XXL v1.1 | 11B | AutoencoderKL | +| Cosmos2 | T5 11B | 11B | Wan VAE | +| Cosmos3 | Cosmos3 reasoner | not specified | Wan/Cosmos VAE | +| DeepFloyd IF | T5 XXL v1.1 | 11B | None | +| ERNIE-Image | ERNIE text encoder | not specified | Flux.2 VAE | +| Flux.1 | CLIP-L/14 + T5 XXL v1.1 | 123M + 11B | AutoencoderKL | +| Flux.2 | Mistral-Small-3.1-24B | 24B | Flux.2 VAE | +| HeartMuLa | None | N/A | HeartCodec tokens | +| HiDream | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 + Llama | 123M + 694M + 11B + not specified | AutoencoderKL | +| Hunyuan Video | Hunyuan LLM | not specified | Hunyuan Video 3D VAE | +| Ideogram 4 | Qwen3-VL-8B-Instruct | 8B | Ideogram AutoEncoder | +| Kandinsky 5.0 Image | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Flux VAE (AutoencoderKL) | +| Kandinsky 5.0 Video | Qwen2.5-VL + CLIP-L/14 | 7B + 123M | Hunyuan Video VAE | +| Kwai Kolors | ChatGLM-6B | 6B | AutoencoderKL | +| Krea2 | Qwen3VL | not specified | Qwen Image VAE | +| LongCat Image | Qwen2.5-VL | 7B | AutoencoderKL | +| LongCat Video | Qwen2.5-VL | 7B | Wan VAE | +| LTX Video | T5 XXL v1.1 | 11B | LTX Video VAE | +| LTX Video 2 | Gemma3 | not specified | LTX Video 2 VAE | +| Lumina2 | Gemma2 | 2B | AutoencoderKL | +| Mage-Flow | Qwen3-VL | not specified | Mage-VAE | +| OmniGen | Integrated OmniGen encoder | not specified | AutoencoderKL | +| PixArt Sigma | T5 XXL v1.1 | 11B | AutoencoderKL | +| Qwen Image | Qwen2.5-VL | 7B | Qwen Image VAE | +| Sana | Gemma2 2B-IT | 2B | Sana AutoencoderDC | +| Sana Video | Gemma 2 | 2B | Wan VAE | +| SD 1.x/2.x (Legacy) | CLIP-L/14 | 123M | AutoencoderKL | +| Stable Diffusion 3 | CLIP-L/14 + CLIP-G/14 + T5 XXL v1.1 | 123M + 694M + 11B | AutoencoderKL | +| Stable Diffusion XL | CLIP-L/14 + CLIP-G/14 | 123M + 694M | AutoencoderKL | +| Stable Cascade (Stage C) | CLIP-ViT-bigG-14 | 694M | Stable Cascade Stage C VAE | +| Wan Video | UMT5 | not specified | Wan VAE | +| Wan S2V | UMT5 | not specified | Wan VAE | +| Z-Image | Qwen3 4B | 4B | AutoencoderKL | +| Z-Image Omni | Qwen3 4B | 4B | AutoencoderKL | +| ZLab I1 | T5Gemma 2B | 2B | AutoencoderKL | + +
-| 模型 | 参数量 | PEFT LoRA | Lycoris | 全秩 | 量化 | 混合精度 | 梯度检查点 | Flow Shift | TwinFlow | Self-Flow | LayerSync | Ref Inputs | ControlNet | Sliders† | 指南 | -| --- | --- | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | :---: | --- | -| PixArt Sigma | 0.6B–0.9B | ✗ | ✓ | ✓ | int8 可选 | bf16 | ✓ | ✗ | ✗ | ✓ | ✓ | ✗ | ✓ | ✓ | [SIGMA.md](quickstart/SIGMA.md) | -| NVLabs Sana | 1.6B–4.8B | ✗ | ✓ | ✓ | int8 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [SANA.md](quickstart/SANA.md) | -| Kwai Kolors | 2.7B | ✓ | ✓ | ✓ | 不推荐 | bf16 | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [KOLORS.md](quickstart/KOLORS.md) | -| Stable Diffusion 3 | 2B–8B | ✓ | ✓ | ✓ | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [SD3.md](quickstart/SD3.md) | -| Flux.1 | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [FLUX.md](quickstart/FLUX.md) | -| Flux.2 | 32B | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [FLUX2.md](quickstart/FLUX2.md) | -| Flux Kontext | 8B–12B | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✓ | ✓ | [FLUX_KONTEXT.md](quickstart/FLUX_KONTEXT.md) | -| Z-Image Turbo | 6B | ✓ | ✗ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZIMAGE.md](quickstart/ZIMAGE.md) | -| Krea2 | - | ✓ | ✗ | ✓* | int8 可选 | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✓ opt | ✗ | ✓ | [KREA2.md](quickstart/KREA2.zh.md) | -| Boogu-Image 0.1 | - | ✓ | ✓ | ✓* | fp8 可选 | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ edit | ✗ | ✓ | [BOOGU_IMAGE.md](quickstart/BOOGU_IMAGE.zh.md) | -| zlab i1 | 3B | ✓ | ✓ | ✓ | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [ZLAB_i1.md](quickstart/ZLAB_i1.zh.md) | -| Ideogram 4 | 9B | ✓ | ✓ | ✓* | fp8 默认,nf4 可选 | bf16 | ✓+ | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [IDEOGRAM4.md](quickstart/IDEOGRAM4.zh.md) | -| ACE-Step | 3.5B | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✓ | [ACE_STEP.md](quickstart/ACE_STEP.md) | -| Chroma 1 | 8.9B | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [CHROMA.md](quickstart/CHROMA.md) | -| Auraflow | 6B | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓+ | ✓ (SLG) | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [AURAFLOW.md](quickstart/AURAFLOW.md) | -| HiDream I1 | 17B (8.5B MoE) | ✓ | ✓ | ✓* | int8/fp8/nf4 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ | ✓ | [HIDREAM.md](quickstart/HIDREAM.md) | -| OmniGen | 3.8B | ✓ | ✓ | ✓ | int8/fp8 可选 | bf16 | ✓ | ✓ | ✗ | ✓ | ✗ | ✗ | ✗ | ✓ | [OMNIGEN.md](quickstart/OMNIGEN.md) | -| Stable Diffusion XL | 2.6B | ✓ | ✓ | ✓ | 不推荐 | bf16 | ✓ | ✗ | ✗ | ✗ | ✓ | ✗ | ✓ | ✓ | [SDXL.md](quickstart/SDXL.md) | -| Lumina2 | 2B | ✓ | ✓ | ✓ | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✗ | ✓ | [LUMINA2.md](quickstart/LUMINA2.md) | -| Cosmos2 | 2B | ✓ | ✓ | ✓ | 不推荐 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [COSMOS2IMAGE.md](quickstart/COSMOS2IMAGE.md) | -| Cosmos3 | 16B-65B | ✓ | ✓ | ✓* | no_change first | bf16 | ✓ | ✓ | ✗ | ✗ | ✗ | audio opt | ✗ | ✓ | [COSMOS3.md](quickstart/COSMOS3.zh.md) | -| LTX Video | ~2.5B | ✓ | ✓ | ✓ | int8/fp8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [LTXVIDEO.md](quickstart/LTXVIDEO.md) | -| LTX Video 2 | 19B | ✓ | ✓ | ✓* | int8/fp8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LTXVIDEO2.md](quickstart/LTXVIDEO2.md) | -| Hunyuan Video 1.5 | 8.3B | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [HUNYUANVIDEO.md](quickstart/HUNYUANVIDEO.md) | -| Wan 2.x | 1.3B–14B | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN.md](quickstart/WAN.md) | -| Wan 2.2 S2V | 14B | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [WAN_S2V.md](quickstart/WAN_S2V.md) | -| Qwen Image | 20B | ✓ | ✓ | ✓* | **必需** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [QWEN_IMAGE.md](quickstart/QWEN_IMAGE.md) | -| Qwen Image Edit | 20B | ✓ | ✓ | ✓* | **必需** (int8/nf4) | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [QWEN_EDIT.md](quickstart/QWEN_EDIT.md) | -| Stable Cascade (C) | 1B, 3.6B prior | ✓ | ✓ | ✓* | 不支持 | fp32 (必需) | ✓ | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ | ✓ | [STABLE_CASCADE_C.md](quickstart/STABLE_CASCADE_C.md) | -| Kandinsky 5.0 Image | 6B (lite) | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✗ | ✓ I2I | ✗ | ✓ | [KANDINSKY5_IMAGE.md](quickstart/KANDINSKY5_IMAGE.md) | -| Kandinsky 5.0 Video | 2B (lite), 19B (pro) | ✓ | ✓ | ✓* | int8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ I2V | ✗ | ✓ | [KANDINSKY5_VIDEO.md](quickstart/KANDINSKY5_VIDEO.md) | -| LongCat-Video | 13.6B | ✓ | ✓ | ✓* | int8/fp8 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ opt | ✗ | ✓ | [LONGCAT_VIDEO.md](quickstart/LONGCAT_VIDEO.md) | -| LongCat-Video Edit | 13.6B | ✓ | ✓ | ✓* | int8/fp8 可选 | bf16 | ✓+ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_VIDEO_EDIT.md](quickstart/LONGCAT_VIDEO_EDIT.md) | -| LongCat-Image | 6B | ✓ | ✓ | ✓* | int8/fp8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ | ✓ | [LONGCAT_IMAGE.md](quickstart/LONGCAT_IMAGE.md) | -| LongCat-Image Edit | 6B | ✓ | ✓ | ✓* | int8/fp8 可选 | bf16 | ✓ | ✓ | ✓ | ✓ | ✓ | ✓ req | ✗ | ✓ | [LONGCAT_EDIT.md](quickstart/LONGCAT_EDIT.md) | - -*✓ = 支持,✓* = 全秩训练需要 DeepSpeed/FSDP2,✗ = 不支持,`✓+` 表示由于 VRAM 压力建议启用检查点。Ref Inputs 仅标记现有参考/编辑/I2V 条件路径;`opt` 表示可选,`req` 表示该编辑/I2V flavour 需要它。TwinFlow ✓ 表示当 `twinflow_enabled=true` 时原生支持(扩散模型需要 `diff2flow_enabled+twinflow_allow_diff2flow`)。Self-Flow ✓ 表示原生支持 `crepa_enabled=true`、`crepa_feature_source=self_flow`、`use_ema=true` 且设置 `crepa_teacher_block_index`。LayerSync ✓ 表示骨干网络暴露 transformer 隐藏状态用于自对齐;✗ 标记没有该缓冲区的 UNet 风格骨干网络。†Sliders 适用于 LoRA 和 LyCORIS(包括全秩 LyCORIS "full")。* - -> 注:Wan 快速开始包含 2.1 + 2.2 阶段预设和时间嵌入切换。Flux Kontext 涵盖基于 Flux.1 构建的编辑工作流程。 - -> 警告:这些快速开始指南是持续更新的文档。随着新模型的发布或训练方案的改进,预计会有不定期更新。 +*✓ = 支持,✓* = 支持但 full-rank training 通常需要 DeepSpeed/FSDP2,✗ = 不支持。Ref Inputs 标记现有 reference/edit/I2V conditioning paths;`opt` 表示可选,`req` 表示 edit/I2V flavour 必需。* +*TwinFlow 在 `twinflow_enabled=true` 时为原生支持;diffusion models 仍需要 `diff2flow_enabled=true` 和 `twinflow_allow_diff2flow=true`。Self-Flow 指 CREPA self-flow support。LayerSync 标记公开 hidden states 用于 alignment 的 backbones。* ### 快速通道:Z-Image Turbo 和 Flux Schnell @@ -76,4 +312,4 @@ - 在 WebUI 中,浏览到您的数据集目录并选择它以查看过滤统计 - 在数据集处理期间检查日志中的统计信息,如:`Sample processing statistics: {'total_processed': 100, 'skipped': {'too_small': 15, ...}}` -详细故障排除请参阅数据加载器文档中的[故障排除-过滤后的数据集](DATALOADER.zh.md#故障排除-过滤后的数据集)。 +详细故障排除请参阅数据加载器文档中的[故障排除-过滤后的数据集](DATALOADER.zh.md)。 diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.es.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.es.md new file mode 100644 index 000000000..4a232d5ef --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.es.md @@ -0,0 +1,1007 @@ +# Checkpointing Segmentado + +El checkpointing segmentado queda entre checkpointing en cada block y no usar checkpointing. + +Usa el backend de activation checkpointing de PyTorch. SimpleTuner ejecuta un grupo contiguo de transformer blocks bajo una llamada de checkpoint y pasa el hidden state devuelto al grupo siguiente. Grupos mas anchos recomputan menos en backward, pero mantienen mas activations vivas. + +Para CPU offload y FFN-only checkpointing, usa [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.es.md#controls). La regla corta sigue en [Decision Rule](UNSLOTH_CHECKPOINTING.es.md#decision-rule). + +## Controles + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +En rutas whole-block soportadas, `gradient_checkpointing_interval` es el ancho del segment. `2` significa checkpoint de blocks `0-1`, `2-3`, `4-5`, etc. + +Para controlar VRAM con mas detalle, agrega stride: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +Esto checkpointa blocks `0-1`, ejecuta `2-3` normalmente, checkpointa `4-5`, ejecuta `6-7` normalmente y repite. El stride debe ser al menos el interval; schedules solapados no son validos. + +Rutas segmented whole-block soportadas: Flux.1, Flux.2, HunyuanVideo, Krea 2, LongCat Image, LongCat Video, LTXVideo 0.9, LTXVideo2, Lumina2, MageFlow, PixArt, SD3, SanaVideo, Z-Image, ZLab I1 y Wan. + +Stable Cascade stage C tambien soporta interval y stride, pero aplica el schedule a la secuencia de micro-bloques Res/Timestep/Attention del UNet en vez de a grupos transformer whole-block. + +Algunas familias usan semantica de interval especifica del modelo: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | Checkpoint cada N-th block | Ignorado | +| Stable Cascade stage C | Checkpoint de micro-bloques UNet por interval | Stride alterna ventanas UNet checkpointed y no checkpointed | +| SD1x, SDXL | Sin soporte segmented whole-block | Ignorado | + +No compares filas stride cuando stride se ignora. Si los numeros son identicos, normalmente el option no se aplico. + +## Cuando Usarlo + +Usalo despues de que el checkpointing normal por block ya entra en VRAM pero cuesta demasiado tiempo por step. Empieza con `2`. Si hay VRAM, prueba `2` con stride `4` en modelos muy profundos. + +No esperes que ayude cuando el peak viene sobre todo de pesos entrenables, optimizer state, validation, cache VAE, block swapping o routing. SimpleTuner vuelve a la ruta per-block mas segura cuando una funcion del modelo necesita control por block. + +`dynamo_use_regional_compilation` no es una victoria universal. Ayudo o fue neutral en varios image models, pero fue mala opcion para los perfiles Wan/RamTorch y LTXVideo2 de abajo. + +## Benchmarks + +Medido con ejemplos reales de SimpleTuner en pods single-GPU H100 y L40S. Validation y checkpoint saves estuvieron desactivados, cache preparation quedo fuera, y el compile/setup del primer step dentro del train loop se excluye cuando existe timing post-warmup. + +Cada celda medida es `post-warmup sec/step / peak VRAM GiB`. Las celdas de estado significan: `OOM` agoto memoria de GPU, `failed` no llego a training steps medidos, `unsupported` significa que esa opcion no estaba conectada para esa familia, y `not run` significa que el sweep no incluyo esa combinacion. + +Compara modos dentro de una familia primero. Las comparaciones entre familias son aproximadas porque cambian resolution, frame count, attention backend, model depth, trainable adapter type y dataset shape. + +La matriz de abajo es la fuente de verdad para este sweep. Las notas especificas por modelo marcan caveats cuando una fila es coverage data y no una recomendacion. + + + +### Resultados Por Familia + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Destilacion AnyFlow en Anima + +Ejemplo: `anima-anyflow.peft-lora`. Resolucion: 1024x1024. Esta fila mide destilacion AnyFlow usando Anima, no el ejemplo LoRA de Anima puro. Usa `anima.peft-lora` para entrenamiento de imagen Anima puro a 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +Nota: AuraFlow soporta cuantizacion con SDNQ y TorchAO. Las filas cuantizadas sin checkpointing estan abajo; las filas cuantizadas con checkpointing necesitan una nueva medicion completa antes de publicar numeros aqui. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. Entrenamiento con tokens de audio. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Ejemplo: `hunyuanvideo-1.5-t2v.peft-lora`. Forma de training: video buckets de 480 pixel-area, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo usa muchas activations con esta forma de training. Per-block e interval-2 encajan bien; sin checkpointing no entra en un H100 de 80 GB. `seg2-stride4` solo entro en este sweep con attention activation offload, y esa fila es una salida de memoria, no una recomendacion de velocidad. SDNQ Hadamard funciona, pero las formas variables de conditioning todavia disparan compilacion de kernels dinamicos durante la medicion. + +### Ideogram 4.0 + +Ejemplo: `ideogram-fp8.peft-lora`. Resolucion: 1024x1024. El flavour fp8 usa el checkpoint fp8 nativo weight-only de Ideogram 4 (`base_model_precision=no_change`). + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +Ejemplo: `krea2.peft-lora`. Resolucion de entrenamiento: recorte cuadrado de 512 px. Ajuste de validacion del ejemplo: 1024x1024; la validacion estuvo desactivada en el benchmark. + +La tabla principal usa regional compilation. Eso ayuda al step speed de Krea2, pero no da una comparacion limpia de VRAM: el compiled graph/workspace mantiene el pico cerca del pico sin checkpointing en varios modos. Una corrida de control bf16 con regional compilation desactivado mostro que el checkpointing esta conectado y tiene la forma esperada de memoria/velocidad: + +La fila `activation-offload` aqui significa full-block checkpointing mas attention activation offload. Frente a full-block `layer` checkpointing solo, attention offload no redujo el peak VRAM de Krea2 en este shape; principalmente agrego CPU transfer overhead. + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +Ejemplo: `longcat-image.peft-lora`. Resolucion de entrenamiento: 512 px cuadrada; resolucion de validacion: 1024x1024. Las filas usan `attention_mechanism=native-flash`. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +Ejemplo: `longcat-video.peft-lora+ramtorch`. Resolucion: 832x480, 81f. Las filas usan `attention_mechanism=native-flash`. + +LongCat Video usa muchas activations con esta forma. Full per-block checkpointing es la fila practica. Las filas partial checkpoint (`interval2`, `seg2-stride4`) no caben aqui, incluso con attention activation offload en la fila strided. Attention activation offload simple cabe para bf16 y SDNQ, pero es mucho mas lento que full checkpointing. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +Los numeros son segundos warm por step / GiB pico. El promedio de la ejecucion completa incluye setup y compile, y queda en los artifacts del sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +Las filas de attention activation offload no estan soportadas para LTXVideo 0.9 en este sweep. + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +Nota: Lumina2 ahora usa la ruta segmented whole-block. `interval2` checkpointa cada segmento de dos blocks; `seg2-stride4` checkpointa dos blocks, deja que los dos siguientes conserven activations y repite. Attention activation offload no fue parte de esta corrida de Lumina2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +Nota: la ruta de imagen variable de MageFlow a 1024px se beneficia sobre todo de attention activation offload y FP8 weight-only. Los modos de block checkpointing son validos, pero no redujeron el pico de residencia medido en este sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +Nota: OmniGen usa prompts como token IDs en vez de embeddings de texto cacheados. Estas filas miden los caminos soportados sin checkpointing y con checkpointing torch de bloque completo; los controles interval, segmented-stride y attention-offload no estan implementados para esta familia en este barrido. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo usa linear attention, asi que attention activation offload sigue unsupported. Segmented whole-block checkpointing esta soportado en la ruta standard. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +Nota: SD1x usa la ruta UNet de diffusers. El checkpointing por capa funciona, pero los controles interval y segmented stride no estan conectados para esta familia. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +Nota: SD3 usa checkpointing segmentado contiguo real en la ruta transformer simple. Attention activation offload esta soportado; reduce mucho la VRAM, pero cuesta throughput. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +Nota: stage C usa la ruta prior en precision completa. Estas filas se ejecutaron con `mixed_precision=no` y `base_model_precision=no_change`; las filas de precision base cuantizada no son significativas para este modelo. Los modos interval y stride operan sobre la secuencia de micro-bloques Res/Timestep/Attention del UNet. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +Nota: ZLab I1 lleva sus skip tensors tipo U-Net dentro del estado segmented checkpoint. Attention activation offload no esta conectado para esta familia. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.hi.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.hi.md new file mode 100644 index 000000000..359e1f1d4 --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.hi.md @@ -0,0 +1,1007 @@ +# Segmented Checkpointing + +Segmented checkpointing हर block को checkpoint करने और कोई checkpointing न करने के बीच का mode है। + +यह PyTorch activation checkpointing backend इस्तेमाल करता है। SimpleTuner contiguous transformer blocks के group को एक checkpoint call में चलाता है, फिर returned hidden state अगले group को देता है। Wider groups backward में कम recompute करते हैं, लेकिन ज्यादा activations alive रखते हैं। + +CPU offload और FFN-only checkpointing के लिए [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.hi.md#controls) देखें। छोटी rule of thumb [Decision Rule](UNSLOTH_CHECKPOINTING.hi.md#decision-rule) में है। + +## Controls + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +Supported whole-block paths पर `gradient_checkpointing_interval` segment width है। `2` का मतलब blocks `0-1`, `2-3`, `4-5` वगैरह checkpoint होंगे। + +ज्यादा fine VRAM control के लिए stride जोड़ें: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +यह blocks `0-1` checkpoint करता है, `2-3` normal चलाता है, `4-5` checkpoint करता है, `6-7` normal चलाता है, और repeat करता है। stride कम से कम interval जितना होना चाहिए; overlapping schedules valid नहीं हैं। + +Supported segmented whole-block paths: Flux.1, Flux.2, HunyuanVideo, Krea 2, LongCat Image, LongCat Video, LTXVideo 0.9, LTXVideo2, Lumina2, MageFlow, PixArt, SD3, SanaVideo, Z-Image, ZLab I1, और Wan। + +Stable Cascade stage C भी interval और stride support करता है, लेकिन schedule transformer whole-block groups की जगह UNet के Res/Timestep/Attention micro-block sequence पर apply होता है। + +कुछ model families पुराने semantics use करते हैं: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | हर N-th block checkpoint | Ignore | +| Stable Cascade stage C | interval से UNet micro-blocks checkpoint | stride checkpointed और non-checkpointed UNet micro-block windows alternate करता है | +| SD1x, SDXL | segmented whole-block support नहीं | Ignore | + +जहां stride ignore होता है, वहां stride rows compare न करें। अगर numbers identical दिखें, तो usually option apply नहीं हुआ। + +## कब इस्तेमाल करें + +जब normal per-block checkpointing fit हो जाए लेकिन step time बहुत महंगा हो, तब इसे इस्तेमाल करें। `2` से शुरू करें। अगर VRAM allow करे, बहुत deep models पर `2` with stride `4` try करें। + +जब peak मुख्य रूप से trainable weights, optimizer state, validation, VAE caching, block swapping, या routing से हो, तो इससे मदद की उम्मीद न करें। जब model feature को per-block control चाहिए, SimpleTuner safer per-block path पर वापस जाता है। + +`dynamo_use_regional_compilation` universal win नहीं है। कुछ image models पर यह helpful या neutral था, लेकिन नीचे Wan/RamTorch और LTXVideo2 profiles में खराब fit था। + +## Benchmarks + +Real SimpleTuner examples से single-GPU H100 और L40S pods पर measure किया गया। validation और checkpoint saves disabled थे, cache preparation exclude था, और post-warmup timing available होने पर train loop के first-step compile/setup को भी exclude किया गया। + +हर measured cell `post-warmup sec/step / peak VRAM GiB` है। status-only cells का मतलब है: `OOM` GPU memory खत्म होना, `failed` measured training steps तक न पहुंचना, `unsupported` उस family में option wired न होना, और `not run` उस combination का sweep में न होना। + +पहले एक ही family के अंदर modes compare करें। Cross-family comparison rough है क्योंकि resolution, frame count, attention backend, model depth, trainable adapter type, और dataset shape अलग हो सकते हैं। + +नीचे की matrix इस sweep की source of truth है। Model-specific notes caveats बताते हैं जब कोई row recommendation के बजाय coverage data हो। + + + +### Family Sweep Results + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Anima पर AnyFlow Distillation + +Example: `anima-anyflow.peft-lora`। Resolution: 1024x1024। यह row Anima के साथ AnyFlow distillation मापती है, plain Anima LoRA example नहीं। plain 1024x1024 Anima image training के लिए `anima.peft-lora` use करें। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +Note: AuraFlow SDNQ और TorchAO quantization support करता है। Quantized `none` rows नीचे हैं; quantized checkpoint rows में numbers डालने से पहले fresh full-length benchmark coverage चाहिए। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. Audio-token training. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Example: `hunyuanvideo-1.5-t2v.peft-lora`. Training shape: 480 pixel-area video buckets, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo इस training shape पर activation-heavy है। Per-block और interval-2 checkpointing साफ चलते हैं; checkpointing बंद करने पर 80 GB H100 भी पर्याप्त नहीं है। `seg2-stride4` इस sweep में attention activation offload के साथ ही चला, इसलिए वह speed recommendation नहीं बल्कि memory fallback है। SDNQ Hadamard काम करता है, लेकिन variable conditioning shapes measurement window में भी dynamic kernel compilation trigger करते हैं। + +### Ideogram 4.0 + +Example: `ideogram-fp8.peft-lora`. Resolution: 1024x1024. `fp8` flavour Ideogram 4 के native weight-only fp8 checkpoint का उपयोग करता है (`base_model_precision=no_change`). + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +Example: `krea2.peft-lora`. Training resolution: 512 px square crop. Example validation setting: 1024x1024; benchmark में validation disabled था. + +Main table regional compilation use करती है। यह Krea2 step speed के लिए अच्छा है, लेकिन VRAM comparison साफ नहीं रहता: compiled graph/workspace कई checkpoint modes में peak को uncheckpointed peak के करीब रखता है। regional compilation बंद करके bf16 control run ने दिखाया कि checkpointing wired है और memory/speed shape expected है: + +यहां `activation-offload` row का मतलब full-block checkpointing plus attention activation offload है। सिर्फ full-block `layer` checkpointing के मुकाबले, इस Krea2 shape में attention offload ने peak VRAM घटाया नहीं; इसने mostly CPU transfer overhead जोड़ा। + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +Example: `longcat-image.peft-lora`. Training resolution: 512 px square; validation resolution: 1024x1024. Rows use `attention_mechanism=native-flash`. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +Example: `longcat-video.peft-lora+ramtorch`. Resolution: 832x480, 81f. Rows `attention_mechanism=native-flash` use करते हैं। + +LongCat Video इस shape पर activation-heavy है। Full per-block checkpointing practical row है। Partial checkpoint rows (`interval2`, `seg2-stride4`) यहाँ fit नहीं होते, strided row पर attention activation offload enabled होने पर भी नहीं। Plain attention activation offload bf16 और SDNQ में fit होता है, लेकिन full checkpointing से काफी धीमा है। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +Numbers warm seconds/step / peak GiB हैं। full-run average setup और compile overhead include करता है, इसलिए वह sweep artifacts में रखा गया है। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +इस sweep में LTXVideo 0.9 के लिए attention activation offload rows supported नहीं हैं। + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +Note: Lumina2 अब segmented whole-block path इस्तेमाल करता है। `interval2` हर two-block segment को checkpoint करता है; `seg2-stride4` दो blocks checkpoint करता है, अगले दो blocks को activations रखने देता है, फिर repeat करता है। Attention activation offload इस Lumina2 run में शामिल नहीं था। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +Note: MageFlow के 1024px variable-shape image path में सबसे ज्यादा फायदा attention activation offload और weight-only FP8 से मिला। Block checkpointing modes valid हैं, लेकिन इस sweep में measured peak residency कम नहीं हुई। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +Note: OmniGen cached text embeddings की जगह token-ID prompts use करता है। ये rows supported no-checkpointing और full-block torch checkpointing paths को measure करती हैं; interval, segmented-stride, और attention-offload controls इस family के लिए इस sweep में implemented नहीं हैं। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo linear attention use करता है, इसलिए attention activation offload unsupported रहता है। Standard path में segmented whole-block checkpointing supported है। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +Note: SD1x diffusers UNet path use करता है. Regular layer checkpointing supported है, लेकिन interval और segmented stride controls इस family में wired नहीं हैं. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +Note: SD3 plain transformer path पर real contiguous segmented checkpointing use करता है. Attention activation offload supported है; यह VRAM काफी घटाता है, लेकिन throughput खर्च करता है. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +Note: stage C full-precision prior path use करता है। ये rows `mixed_precision=no` और `base_model_precision=no_change` के साथ run हुईं; quantized base precision rows इस model के लिए meaningful नहीं हैं। interval और stride modes UNet के Res/Timestep/Attention micro-block sequence पर काम करते हैं। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +Note: ZLab I1 अपने U-Net-style skip tensors को segmented checkpoint state में carry करता है। इस family में attention activation offload wired नहीं है। + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.ja.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.ja.md new file mode 100644 index 000000000..b53b9fbcf --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.ja.md @@ -0,0 +1,1007 @@ +# Segmented Checkpointing + +Segmented checkpointing は、全 block を checkpoint する設定と checkpoint なしの中間にある設定です。 + +PyTorch の activation checkpointing backend を使います。SimpleTuner は連続した transformer blocks を 1 回の checkpoint 呼び出しで実行し、返された hidden state を次の group に渡します。group を広くすると backward の再計算は減りますが、保持する activations は増えます。 + +CPU offload と FFN-only checkpointing は [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.ja.md#controls) を参照してください。短い判断基準は [Decision Rule](UNSLOTH_CHECKPOINTING.ja.md#decision-rule) にあります。 + +## Controls + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +対応している whole-block path では、`gradient_checkpointing_interval` が segment width です。`2` は blocks `0-1`、`2-3`、`4-5` を checkpoint します。 + +より細かく VRAM を制御するには stride を追加します: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +これは blocks `0-1` を checkpoint し、`2-3` は通常実行、`4-5` を checkpoint、`6-7` は通常実行、という形で繰り返します。stride は interval 以上である必要があり、重複 schedule は無効です。 + +Segmented whole-block path 対応: Flux.1、Flux.2、HunyuanVideo、Krea 2、LongCat Image、LongCat Video、LTXVideo 0.9、LTXVideo2、Lumina2、MageFlow、PixArt、SD3、SanaVideo、Z-Image、ZLab I1、Wan。 + +Stable Cascade stage C も interval と stride に対応していますが、schedule は transformer whole-block group ではなく UNet の Res/Timestep/Attention micro-block sequence に適用されます。 + +一部の family は古い semantics です: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | N block ごとに checkpoint | 無視 | +| Stable Cascade stage C | interval ごとに UNet micro-block を checkpoint | stride は checkpointed/non-checkpointed UNet micro-block window を交互にする | +| SD1x, SDXL | segmented whole-block 非対応 | 無視 | + +stride が無視される family では stride 行を比較しないでください。数字が同じなら、多くの場合は option が効いていないだけです。 + +## 使う場面 + +通常の per-block checkpointing では fit するが step time が高すぎる場合に使います。まず `2` から始めてください。VRAM に余裕がある深いモデルでは `2` と stride `4` を試します。 + +peak の主因が trainable weights、optimizer state、validation、VAE caching、block swapping、routing の場合は効果を期待しないでください。モデル機能が per-block 制御を必要とする場合、SimpleTuner はより安全な per-block path に戻ります。 + +`dynamo_use_regional_compilation` は万能ではありません。一部の image model では有利または中立でしたが、下の Wan/RamTorch と LTXVideo2 profile では悪い結果でした。compile 設定も benchmark 条件として扱ってください。 + +## Benchmarks + +実際の SimpleTuner examples を使い、single-GPU H100/L40S pods で測定しました。validation と checkpoint saves は無効、cache preparation は除外、post-warmup timing がある場合は train loop 内の first-step compile/setup も除外しています。 + +測定済み cell は `post-warmup sec/step / peak VRAM GiB` です。status-only cell の意味は、`OOM` が GPU memory 不足、`failed` が measured training steps まで到達しなかった run、`unsupported` がその family に option が wired されていない状態、`not run` が sweep にその組み合わせがなかった状態です。 + +まず同じ family 内で mode を比較してください。family 間の比較は、resolution、frame count、attention backend、model depth、trainable adapter type、dataset shape が違うため、おおまかな目安です。 + +下の matrix がこの sweep の source of truth です。model-specific notes は、row が recommendation ではなく coverage data の場合に caveat を示します。 + + + +### Family Sweep Results + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Anima 上の AnyFlow 蒸留 + +例: `anima-anyflow.peft-lora`。解像度: 1024x1024。この行は Anima を使った AnyFlow 蒸留の測定で、通常の Anima LoRA 例ではありません。通常の 1024x1024 Anima 画像 training には `anima.peft-lora` を使ってください。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +注記: AuraFlow は SDNQ と TorchAO 量子化に対応しています。量子化した `none` 行は下にあります。量子化 checkpoint 行は、新しい full-length benchmark で確認してから数値を入れます。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. Audio-token training. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Example: `hunyuanvideo-1.5-t2v.peft-lora`. Training shape: 480 pixel-area video buckets, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo はこの training shape では activation が重いです。Per-block と interval-2 checkpointing は安定して入り、checkpointing なしは 80 GB H100 でも入りません。`seg2-stride4` は attention activation offload を有効にした場合だけ入りましたが、速度向けではなくメモリ用の fallback です。SDNQ Hadamard は動きますが、conditioning shape が変わるため測定中にも dynamic kernel compilation が残ります。 + +### Ideogram 4.0 + +Example: `ideogram-fp8.peft-lora`. Resolution: 1024x1024. `fp8` flavour は Ideogram 4 の native weight-only fp8 checkpoint を使います (`base_model_precision=no_change`)。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +例: `krea2.peft-lora`。学習解像度: 512 px square crop。例の validation 設定: 1024x1024。benchmark では validation は無効です。 + +メイン表は regional compilation を使っています。Krea2 の step speed には有利ですが、VRAM 比較としてはきれいではありません。compiled graph/workspace により、いくつかの checkpoint mode で peak が未 checkpoint の peak に近く残ります。regional compilation を無効にした bf16 control run では、checkpointing が正しく接続され、期待通りの memory/speed 形になりました: + +ここでの `activation-offload` 行は、full-block checkpointing に attention activation offload を足したものです。full-block `layer` checkpointing 単体と比べると、この Krea2 shape では attention offload は peak VRAM を下げず、主に CPU transfer overhead を増やしました。 + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +例: `longcat-image.peft-lora`。Training resolution: 512 px square; validation resolution: 1024x1024. Rows use `attention_mechanism=native-flash`. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +例: `longcat-video.peft-lora+ramtorch`。Resolution: 832x480, 81f。各行は `attention_mechanism=native-flash` を使用。 + +LongCat Video はこの shape では activation が重いです。Full per-block checkpointing が実用的な行です。Partial checkpoint 行(`interval2`、`seg2-stride4`)はここでは入りません。strided 行に attention activation offload を足しても不足します。通常の attention activation offload は bf16 と SDNQ では入りますが、full checkpointing よりかなり遅くなります。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +数字は warm seconds/step / peak GiB です。full-run average は setup と compile overhead を含むため、sweep artifacts に記録しています。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +Attention activation offload rows は、この sweep の LTXVideo 0.9 では非対応です。 + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +注記: Lumina2 は segmented whole-block path を使うようになりました。`interval2` は 2-block segment ごとに checkpoint し、`seg2-stride4` は 2 blocks を checkpoint、次の 2 blocks は activations を保持して実行し、それを繰り返します。Attention activation offload はこの Lumina2 run には含めていません。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +Note: MageFlow の 1024px variable-shape image path では、主に attention activation offload と weight-only FP8 が効きました。Block checkpointing mode は有効ですが、この sweep では measured peak residency は下がりませんでした。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +Note: OmniGen は cached text embeddings ではなく token-ID prompt を使います。この行は、対応済みの no-checkpointing と full-block torch checkpointing path を測定したものです。interval、segmented-stride、attention-offload controls は、この family ではこの sweep 時点で未実装です。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo は linear attention を使うため、attention activation offload は未対応です。標準 path では segmented whole-block checkpointing に対応しています。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +Note: SD1x は diffusers UNet path を使います。通常の layer checkpointing は supported ですが、interval と segmented stride controls はこの family では未接続です。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +Note: SD3 は plain transformer path で real contiguous segmented checkpointing を使います。Attention activation offload も supported です。VRAM は大きく下がりますが、throughput は落ちます。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +注記: stage C は full precision の prior path です。これらの行は `mixed_precision=no` と `base_model_precision=no_change` で実行しました。量子化 base precision 行はこのモデルでは有用な測定ではありません。interval と stride は UNet の Res/Timestep/Attention micro-block sequence に適用されます。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +注記: ZLab I1 は U-Net 風の skip tensors を segmented checkpoint state に入れて運びます。この family では attention activation offload は未接続です。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.md new file mode 100644 index 000000000..ff697a741 --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.md @@ -0,0 +1,1007 @@ +# Segmented Checkpointing + +Segmented checkpointing is the middle gear between checkpointing every block and checkpointing nothing. + +It uses PyTorch's activation checkpointing backend. SimpleTuner runs a contiguous group of transformer blocks under one checkpoint call, then carries the returned hidden state into the next group. Wider groups recompute less in backward, but keep more activations alive. + +For CPU offload and FFN-only checkpointing, use [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.md#controls). The short rule of thumb still lives in [Decision Rule](UNSLOTH_CHECKPOINTING.md#decision-rule). + +## Controls + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +On supported whole-block paths, `gradient_checkpointing_interval` is the segment width. `2` means checkpoint blocks `0-1`, `2-3`, `4-5`, and so on. + +For finer VRAM control, add a stride: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +That checkpoints blocks `0-1`, runs `2-3` normally, checkpoints `4-5`, runs `6-7` normally, and repeats. The stride must be at least the interval; overlapping schedules are not valid. + +Supported segmented whole-block paths: Flux.1, Flux.2, HunyuanVideo, Krea 2, LongCat Image, LongCat Video, LTXVideo 0.9, LTXVideo2, Lumina2, MageFlow, PixArt, SD3, SanaVideo, Z-Image, ZLab I1, and Wan. + +Stable Cascade stage C also supports interval and stride control, but it applies the schedule to the UNet Res/Timestep/Attention micro-block sequence instead of transformer whole-block groups. + +Some model families use model-specific interval semantics: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | Checkpoint every N-th block | Ignored | +| Stable Cascade stage C | Checkpoint UNet micro-blocks by interval | Stride alternates checkpointed and non-checkpointed UNet micro-block windows | +| SD1x, SDXL | No segmented whole-block support | Ignored | + +Do not compare stride rows for families where stride is ignored. If the benchmark numbers look identical there, that is usually the option being ignored, not a useful performance result. + +## When To Use It + +Use it after normal per-block checkpointing fits but costs too much step time. Start with `2`. If VRAM allows, try `2` with stride `4` on very deep models. + +Do not expect it to help when the peak is mostly trainable weights, optimizer state, validation, VAE caching, block swapping, or routing. SimpleTuner falls back to the safer per-block path when a model feature needs per-block control. + +`dynamo_use_regional_compilation` is not a universal win. It helped or stayed neutral on several image-model runs, but it was a bad fit for the Wan/RamTorch and LTXVideo2 profiles below. Treat compile settings as part of the benchmark, not as background noise. + +## Benchmarks + +Measured with real SimpleTuner examples on single-GPU H100 and L40S pods. Validation and checkpoint saves were disabled, cache preparation was excluded, and first-step compile/setup inside the train loop is excluded when post-warmup timing is available. + +Each measured cell is `post-warmup sec/step / peak VRAM GiB`. Status-only cells mean: `OOM` ran out of GPU memory, `failed` did not reach measured training steps, `unsupported` means that option was not wired for that family, and `not run` means the sweep did not include that combination. + +Compare modes within a family first. Cross-family comparisons are rough because resolution, frame count, attention backend, model depth, trainable adapter type, and dataset shape differ. + +The matrix below is the source of truth for this sweep. Model-specific notes call out caveats when a row is coverage data rather than a recommendation. + + + +### Family Sweep Results + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AnyFlow Distillation on Anima + +Example: `anima-anyflow.peft-lora`. Resolution: 1024x1024. This row measures AnyFlow distillation using Anima, not the plain Anima LoRA example. Use `anima.peft-lora` for plain 1024x1024 Anima image training. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +Note: AuraFlow supports SDNQ and TorchAO quantization. Quantized `none` rows are included below; quantized checkpoint rows need fresh full-length benchmark coverage before they get numbers here. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. Audio-token training. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Example: `hunyuanvideo-1.5-t2v.peft-lora`. Training shape: 480 pixel-area video buckets, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo is activation-heavy at this training shape. Per-block and interval-2 checkpointing both fit cleanly; leaving checkpointing off does not fit on an 80 GB H100. `seg2-stride4` only fit in this sweep when attention activation offload was enabled, and that row is a fallback rather than a speed recommendation. SDNQ Hadamard works, but variable conditioning shapes still trigger dynamic-kernel compilation in the measured window. + +### Ideogram 4.0 + +Example: `ideogram-fp8.peft-lora`. Resolution: 1024x1024. The fp8 flavour uses Ideogram 4's native weight-only fp8 checkpoint (`base_model_precision=no_change`). + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +Example: `krea2.peft-lora`. Training resolution: 512 px square crop. Example validation setting: 1024x1024; validation was disabled for the benchmark. + +The main table uses regional compilation, which is good for Krea2 step speed but not a clean VRAM comparison: the compiled graph/workspace keeps the peak close to the uncheckpointed peak for several checkpoint modes. A bf16 control run with regional compilation disabled showed checkpointing is wired and has the expected memory/speed shape: + +The `activation-offload` row here means full-block checkpointing plus attention activation offload. Against full-block `layer` checkpointing alone, attention offload did not reduce Krea2 peak VRAM in this shape; it mostly added CPU transfer overhead. + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +Example: `longcat-image.peft-lora`. Training resolution: 512 px square; validation resolution: 1024x1024. Rows use `attention_mechanism=native-flash`. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +Example: `longcat-video.peft-lora+ramtorch`. Resolution: 832x480, 81f. Rows use `attention_mechanism=native-flash`. + +LongCat Video is activation-heavy at this shape. Full per-block checkpointing is the practical row. The partial checkpoint rows (`interval2`, `seg2-stride4`) do not fit here, even when attention activation offload is enabled for the strided row. Plain attention activation offload fits for bf16 and SDNQ, but it is much slower than full checkpointing. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +Numbers are warm seconds per step / peak GiB. The full-run average includes setup and compile overhead and is recorded in the sweep artifacts. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +Attention activation offload rows are unsupported for LTXVideo 0.9 in this sweep. + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +Note: Lumina2 now uses the segmented whole-block path. `interval2` checkpoints every two-block segment; `seg2-stride4` checkpoints two blocks, lets the next two keep activations, then repeats. Attention activation offload was not part of this Lumina2 run. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +Note: MageFlow's 1024px variable-shape image path is mostly helped by attention activation offload and weight-only FP8. Block checkpointing modes are valid, but did not lower measured peak residency in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +Note: OmniGen uses token-ID prompts instead of cached text embeddings. These rows measure the supported no-checkpoint and full-block torch checkpointing paths; interval, segmented-stride, and attention-offload controls are not implemented for this family in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo uses linear attention, so attention activation offload remains unsupported. Segmented whole-block checkpointing is supported for the standard path. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +Note: SD1x uses the diffusers UNet path. Regular layer checkpointing is supported, but interval and segmented stride controls are not wired for this family. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +Note: SD3 uses true contiguous segmented checkpointing on the plain transformer path. Attention activation offload is supported; it cuts VRAM hard, but costs throughput. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +Note: Stage C is a full-precision prior path. These rows ran with `mixed_precision=no` and `base_model_precision=no_change`; quantized base precision rows are not meaningful for this model. The interval and stride modes operate over the UNet's Res/Timestep/Attention micro-block sequence. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +Note: ZLab I1 carries its U-Net-style skip tensors through the segmented checkpoint state. Attention activation offload is not wired for this family. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.pt-BR.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.pt-BR.md new file mode 100644 index 000000000..da5b4cf20 --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.pt-BR.md @@ -0,0 +1,1007 @@ +# Checkpointing Segmentado + +Checkpointing segmentado e o meio-termo entre checkpoint em todo block e nenhum checkpoint. + +Ele usa o backend de activation checkpointing do PyTorch. O SimpleTuner executa um grupo contiguo de transformer blocks em uma chamada de checkpoint e passa o hidden state retornado para o proximo grupo. Grupos mais largos recomputam menos no backward, mas mantem mais activations vivas. + +Para CPU offload e FFN-only checkpointing, veja [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.pt-BR.md#controls). A regra curta continua em [Decision Rule](UNSLOTH_CHECKPOINTING.pt-BR.md#decision-rule). + +## Controles + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +Nos caminhos whole-block suportados, `gradient_checkpointing_interval` e a largura do segment. `2` significa checkpoint dos blocks `0-1`, `2-3`, `4-5` e assim por diante. + +Para controle mais fino de VRAM, adicione stride: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +Isso faz checkpoint dos blocks `0-1`, executa `2-3` normalmente, faz checkpoint de `4-5`, executa `6-7` normalmente e repete. O stride deve ser pelo menos o interval; schedules sobrepostos nao sao validos. + +Caminhos segmented whole-block suportados: Flux.1, Flux.2, HunyuanVideo, Krea 2, LongCat Image, LongCat Video, LTXVideo 0.9, LTXVideo2, Lumina2, MageFlow, PixArt, SD3, SanaVideo, Z-Image, ZLab I1 e Wan. + +Stable Cascade stage C tambem suporta interval e stride, mas aplica o schedule a sequencia de micro-blocos Res/Timestep/Attention do UNet em vez de grupos transformer whole-block. + +Algumas familias usam semantica de interval especifica do modelo: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | Checkpoint a cada N-th block | Ignorado | +| Stable Cascade stage C | Checkpoint de micro-blocos UNet por interval | Stride alterna janelas UNet checkpointed e non-checkpointed | +| SD1x, SDXL | Sem suporte segmented whole-block | Ignorado | + +Nao compare linhas stride quando stride e ignorado. Se os numeros forem identicos, normalmente a option nao foi aplicada. + +## Quando Usar + +Use depois que o checkpointing normal por block couber, mas custar tempo demais por step. Comece com `2`. Se houver VRAM, tente `2` com stride `4` em modelos muito profundos. + +Nao espere ajuda quando o peak vem principalmente de pesos treinaveis, optimizer state, validation, cache de VAE, block swapping ou routing. O SimpleTuner volta para o caminho per-block mais seguro quando uma feature do modelo precisa desse controle. + +`dynamo_use_regional_compilation` nao e ganho universal. Ele ajudou ou foi neutro em alguns image models, mas foi ruim nos perfis Wan/RamTorch e LTXVideo2 abaixo. + +## Benchmarks + +Medido com exemplos reais do SimpleTuner em pods single-GPU H100 e L40S. Validation e checkpoint saves ficaram desativados, cache preparation foi excluida, e o compile/setup do primeiro step dentro do train loop e excluido quando ha timing post-warmup. + +Cada celula medida e `post-warmup sec/step / peak VRAM GiB`. Celulas somente com status significam: `OOM` ficou sem memoria de GPU, `failed` nao chegou aos training steps medidos, `unsupported` significa que a opcao nao estava conectada para essa familia, e `not run` significa que o sweep nao incluiu essa combinacao. + +Compare modos dentro da mesma familia primeiro. Comparacoes entre familias sao aproximadas porque resolution, frame count, attention backend, model depth, trainable adapter type e dataset shape mudam. + +A matriz abaixo e a fonte de verdade para este sweep. Notas especificas por modelo marcam caveats quando uma linha e coverage data em vez de recomendacao. + + + +### Resultados Por Familia + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Destilacao AnyFlow no Anima + +Exemplo: `anima-anyflow.peft-lora`. Resolucao: 1024x1024. Esta linha mede destilacao AnyFlow usando Anima, nao o exemplo LoRA de Anima puro. Use `anima.peft-lora` para treinamento de imagem Anima puro em 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +Nota: AuraFlow suporta quantizacao com SDNQ e TorchAO. As linhas quantizadas sem checkpointing estao abaixo; as linhas quantizadas com checkpointing precisam de uma nova medicao completa antes de receber numeros aqui. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. Treinamento com tokens de áudio. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Exemplo: `hunyuanvideo-1.5-t2v.peft-lora`. Forma de training: video buckets de 480 pixel-area, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo usa muitas activations nessa forma de training. Per-block e interval-2 cabem bem; sem checkpointing nao cabe em um H100 de 80 GB. `seg2-stride4` so coube neste sweep com attention activation offload, e essa linha e uma saida para memoria, nao uma recomendacao de velocidade. SDNQ Hadamard funciona, mas formas variaveis de conditioning ainda disparam compilacao de kernels dinamicos durante a medicao. + +### Ideogram 4.0 + +Exemplo: `ideogram-fp8.peft-lora`. Resolucao: 1024x1024. O flavour fp8 usa o checkpoint fp8 nativo weight-only do Ideogram 4 (`base_model_precision=no_change`). + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +Exemplo: `krea2.peft-lora`. Resolucao de treino: crop quadrado de 512 px. Configuracao de validacao do exemplo: 1024x1024; a validacao ficou desativada no benchmark. + +A tabela principal usa regional compilation. Isso ajuda o step speed do Krea2, mas nao e uma comparacao limpa de VRAM: o compiled graph/workspace mantem o pico perto do pico sem checkpointing em varios modos. Uma execucao de controle bf16 com regional compilation desativado mostrou que o checkpointing esta conectado e tem o formato esperado de memoria/velocidade: + +A linha `activation-offload` aqui significa full-block checkpointing mais attention activation offload. Contra full-block `layer` checkpointing sozinho, attention offload nao reduziu o peak VRAM do Krea2 neste shape; principalmente adicionou CPU transfer overhead. + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +Exemplo: `longcat-image.peft-lora`. Resolucao de treino: 512 px quadrada; resolucao de validacao: 1024x1024. As linhas usam `attention_mechanism=native-flash`. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +Exemplo: `longcat-video.peft-lora+ramtorch`. Resolucao: 832x480, 81f. As linhas usam `attention_mechanism=native-flash`. + +LongCat Video usa muitas activations neste shape. Full per-block checkpointing e a linha pratica. As linhas partial checkpoint (`interval2`, `seg2-stride4`) nao cabem aqui, mesmo com attention activation offload na linha strided. Attention activation offload simples cabe para bf16 e SDNQ, mas e muito mais lento que full checkpointing. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +Os numeros sao segundos warm por step / GiB pico. A media da execucao completa inclui setup e compile, e fica nos artifacts do sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +As linhas de attention activation offload nao sao suportadas para LTXVideo 0.9 neste sweep. + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +Nota: Lumina2 agora usa o caminho segmented whole-block. `interval2` faz checkpoint de cada segmento de dois blocks; `seg2-stride4` checkpointa dois blocks, deixa os dois seguintes manterem activations e repete. Attention activation offload nao fez parte desta rodada de Lumina2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +Nota: o caminho de imagem com shapes variaveis do MageFlow em 1024px se beneficia principalmente de attention activation offload e FP8 weight-only. Os modos de block checkpointing sao validos, mas nao reduziram o pico de residencia medido neste sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +Nota: OmniGen usa prompts como token IDs em vez de embeddings de texto em cache. Estas linhas medem os caminhos suportados sem checkpointing e com checkpointing torch de bloco completo; os controles interval, segmented-stride e attention-offload nao estao implementados para esta familia neste sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo usa linear attention, entao attention activation offload continua unsupported. Segmented whole-block checkpointing e suportado no caminho standard. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +Nota: SD1x usa o caminho UNet do diffusers. O checkpointing por camada funciona, mas os controles interval e segmented stride nao estao conectados para esta familia. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +Nota: SD3 usa checkpointing segmentado contiguo real no caminho transformer simples. Attention activation offload e suportado; reduz bastante a VRAM, mas custa throughput. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +Nota: stage C usa o prior em precision completa. Estas linhas foram executadas com `mixed_precision=no` e `base_model_precision=no_change`; linhas de precision base quantizada nao sao significativas para este modelo. Os modos interval e stride operam sobre a sequencia de micro-blocos Res/Timestep/Attention do UNet. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +Nota: ZLab I1 carrega seus skip tensors estilo U-Net pelo estado segmented checkpoint. Attention activation offload nao esta conectado para esta familia. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/SEGMENTED_CHECKPOINTING.zh.md b/documentation/experimental/SEGMENTED_CHECKPOINTING.zh.md new file mode 100644 index 000000000..020dfd52b --- /dev/null +++ b/documentation/experimental/SEGMENTED_CHECKPOINTING.zh.md @@ -0,0 +1,1007 @@ +# 分段 Checkpointing + +分段 checkpointing 介于每个 block 都 checkpoint 和完全不 checkpoint 之间。 + +它使用 PyTorch activation checkpointing backend。SimpleTuner 会把连续的 transformer blocks 放进一次 checkpoint 调用,然后把返回的 hidden state 传给下一组。组越宽,backward 里的重算越少,但保留的 activations 越多。 + +CPU offload 和 FFN-only checkpointing 请看 [Unsloth-style checkpointing](UNSLOTH_CHECKPOINTING.zh.md#controls)。简短判断规则仍在 [Decision Rule](UNSLOTH_CHECKPOINTING.zh.md#decision-rule)。 + +## 控制项 + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2 +} +``` + +在支持 whole-block 的路径上,`gradient_checkpointing_interval` 是 segment 宽度。`2` 表示 checkpoint blocks `0-1`、`2-3`、`4-5`,依此类推。 + +如果需要更细的 VRAM 控制,可以加 stride: + +```json +{ + "gradient_checkpointing": true, + "gradient_checkpointing_backend": "torch", + "gradient_checkpointing_interval": 2, + "gradient_checkpointing_segment_stride": 4 +} +``` + +这会 checkpoint blocks `0-1`,正常运行 `2-3`,checkpoint `4-5`,正常运行 `6-7`,然后重复。stride 必须大于等于 interval;不支持重叠 schedule。 + +支持 segmented whole-block 的路径:Flux.1、Flux.2、HunyuanVideo、Krea 2、LongCat Image、LongCat Video、LTXVideo 0.9、LTXVideo2、Lumina2、MageFlow、PixArt、SD3、SanaVideo、Z-Image、ZLab I1 和 Wan。 + +Stable Cascade stage C 也支持 interval 和 stride,但 schedule 作用在 UNet 的 Res/Timestep/Attention micro-block 序列上,而不是 transformer whole-block group。 + +一些 model family 使用旧语义: + +| Family | `gradient_checkpointing_interval` | `gradient_checkpointing_segment_stride` | +| --- | --- | --- | +| Sana | 每 N 个 block checkpoint 一次 | 忽略 | +| Stable Cascade stage C | 按 interval checkpoint UNet micro-block | stride 在 checkpointed 和非 checkpointed UNet micro-block window 之间交替 | +| SD1x, SDXL | 不支持 segmented whole-block | 忽略 | + +不要比较 stride 被忽略的 family 的 stride 行。如果数字完全一样,通常是 option 没生效,不是有用的性能结果。 + +## 什么时候使用 + +当普通 per-block checkpointing 能 fit 但 step time 太高时使用。先从 `2` 开始。如果 VRAM 允许,在很深的模型上试 `2` + stride `4`。 + +当 peak 主要来自可训练权重、optimizer state、validation、VAE cache、block swapping 或 routing 时,不要期待它有帮助。当模型特性需要 per-block 控制时,SimpleTuner 会回到更稳的 per-block 路径。 + +`dynamo_use_regional_compilation` 不是通用加速。它在一些 image model 上有帮助或基本中性,但在下面的 Wan/RamTorch 和 LTXVideo2 profile 里效果很差。compile 设置也要当作 benchmark 条件。 + +## Benchmarks + +使用真实 SimpleTuner examples,在单卡 H100 和 L40S pods 上测量。validation 和 checkpoint saves 已关闭,cache preparation 不计入;存在 post-warmup timing 时,也会排除 train loop 内第一个 step 的 compile/setup 时间。 + +每个实测单元格都是 `post-warmup sec/step / peak VRAM GiB`。只有状态的单元格含义是:`OOM` 表示 GPU 显存不足,`failed` 表示没有到达可统计的训练 step,`unsupported` 表示该 family 没有接入这个选项,`not run` 表示 sweep 没有覆盖这个组合。 + +优先在同一个 family 内比较不同 mode。跨 family 比较只能作为粗略参考,因为 resolution、frame count、attention backend、model depth、trainable adapter type 和 dataset shape 都可能不同。 + +下面的矩阵是这个 sweep 的唯一数据源。model-specific notes 会说明某些行只是覆盖情况数据,而不是推荐配置。 + + + +### Family Sweep Results + +### ACE Step 1.5 + +Example: `ace_step-v1-5.peft-lora`. Resolution: 512. + +Note: This sweep did not produce a usable ACE Step throughput row. The status-only entries below should be treated as coverage gaps, not as a recommendation. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | OOM | OOM | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | OOM | OOM | +| int8-sdnq-hadamard | interval2 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### 基于 Anima 的 AnyFlow 蒸馏 + +示例:`anima-anyflow.peft-lora`。分辨率:1024x1024。此行测量的是使用 Anima 的 AnyFlow 蒸馏,不是普通 Anima LoRA 示例。普通 1024x1024 Anima 图像训练请使用 `anima.peft-lora`。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.022 / 17.26 | 0.719 / 17.21 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.232 / 5.60 | 0.903 / 5.56 | +| bf16 | interval2 | 1.252 / 5.60 | 0.897 / 5.56 | +| bf16 | seg2-stride4 | 1.244 / 5.60 | 0.898 / 5.56 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 5.417 / 18.61 | 4.974 / 18.57 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.562 / 4.36 | 4.019 / 4.31 | +| int8-sdnq-hadamard | interval2 | 3.723 / 4.36 | 3.196 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4 | 3.658 / 4.36 | 3.140 / 4.31 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.242 / 45.71 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.810 / 5.51 | 2.576 / 5.46 | +| fp8-torchao | interval2 | 2.846 / 5.51 | 2.581 / 5.46 | +| fp8-torchao | seg2-stride4 | 2.766 / 5.51 | 2.567 / 5.46 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### AuraFlow + +Example: `auraflow.peft-lora`. Resolution: 1024x1024. + +说明:AuraFlow 支持 SDNQ 和 TorchAO 量化。下表包含量化的 `none` 行;量化 checkpoint 行需要新的完整 benchmark 覆盖后再填写数值。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.180 / 19.19 | 0.233 / 19.12 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.824 / 13.37 | 0.833 / 13.32 | +| bf16 | interval2 | 1.764 / 16.21 | 0.877 / 16.14 | +| bf16 | seg2-stride4 | 1.771 / 16.21 | 0.887 / 16.14 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.642 / 12.88 | 0.610 / 12.87 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | not run | not run | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.621 / 24.53 | 0.757 / 24.44 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | not run | not run | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Boogu Image + +Example: `boogu-image-v0.1.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.694 / 59.14 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.907 / 23.44 | 2.641 / 23.39 | +| bf16 | interval2 | 0.912 / 23.44 | 2.649 / 23.39 | +| bf16 | seg2-stride4 | 0.911 / 23.44 | 2.648 / 23.39 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.488 / 53.24 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.878 / 15.20 | 3.309 / 15.15 | +| int8-sdnq-hadamard | interval2 | 1.630 / 34.12 | 2.577 / 34.07 | +| int8-sdnq-hadamard | seg2-stride4 | 1.656 / 34.11 | 2.574 / 34.06 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.713 / 18.48 | 3.778 / 18.44 | +| fp8-torchao | interval2 | 1.731 / 18.48 | 3.777 / 18.44 | +| fp8-torchao | seg2-stride4 | 1.721 / 18.48 | 3.779 / 18.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Chroma + +Example: `chroma.peft-lora`. Resolution: 1024x1024. + +Note: Checkpointed Chroma rows use `attention_mechanism=native-efficient`, which was the stable attention path for this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.454 / 26.18 | 0.559 / 26.13 | +| bf16 | activation-offload | 4.873 / 18.74 | 4.793 / 18.69 | +| bf16 | layer | 1.276 / 17.67 | 1.430 / 17.63 | +| bf16 | interval2 | 1.204 / 21.80 | 1.349 / 21.75 | +| bf16 | seg2-stride4 | 1.200 / 21.78 | 1.382 / 21.74 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.083 / 21.42 | 1.061 / 21.37 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.714 / 10.41 | 1.646 / 10.36 | +| int8-sdnq-hadamard | interval2 | 1.443 / 15.72 | 1.391 / 15.68 | +| int8-sdnq-hadamard | seg2-stride4 | 1.428 / 15.71 | 1.323 / 15.67 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.122 / 45.44 | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.821 / 10.64 | 2.104 / 10.60 | +| fp8-torchao | interval2 | 1.447 / 27.49 | 1.871 / 27.44 | +| fp8-torchao | seg2-stride4 | 1.431 / 27.49 | 1.877 / 27.44 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 2 Image + +Example: `cosmos2image.lycoris-lokr`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.336 / 8.05 | 0.316 / 8.00 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.567 / 4.09 | 0.559 / 4.04 | +| bf16 | interval2 | 0.595 / 4.09 | 0.544 / 4.04 | +| bf16 | seg2-stride4 | 0.598 / 4.09 | 0.546 / 4.04 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.831 / 6.56 | 0.783 / 6.56 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.315 / 2.44 | 1.288 / 2.39 | +| int8-sdnq-hadamard | interval2 | 1.346 / 2.44 | 1.321 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4 | 1.413 / 2.44 | 1.274 / 2.39 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.850 / 13.30 | 0.840 / 13.25 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.692 / 2.70 | 1.560 / 2.65 | +| fp8-torchao | interval2 | 1.626 / 2.70 | 1.544 / 2.65 | +| fp8-torchao | seg2-stride4 | 1.607 / 2.70 | 1.555 / 2.65 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Cosmos 3 + +Example: `cosmos3-edge-image-24g.lycoris-lokr`. Resolution: 1024 px. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 2.747 / 8.90 | 2.904 / 8.86 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.602 / 8.90 | 2.606 / 8.86 | +| bf16 | interval2 | 2.658 / 8.90 | 2.567 / 8.86 | +| bf16 | seg2-stride4 | 2.628 / 8.90 | 2.899 / 8.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 17.112 / 6.21 | 17.965 / 6.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 3.238 / 6.21 | 2.891 / 6.16 | +| int8-sdnq-hadamard | interval2 | 3.253 / 6.21 | 2.916 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4 | 3.254 / 6.21 | 2.923 / 6.16 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.849 / 17.84 | 1.559 / 17.80 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 1.869 / 17.84 | 1.599 / 17.79 | +| fp8-torchao | interval2 | 1.855 / 17.84 | 1.600 / 17.80 | +| fp8-torchao | seg2-stride4 | 1.859 / 17.84 | 1.523 / 17.80 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### ERNIE 4.5 Image + +Example: `ernie.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.711 / 15.61 | 1.282 / 15.56 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.110 / 4.94 | 1.840 / 4.90 | +| bf16 | interval2 | 0.876 / 10.16 | 1.536 / 10.12 | +| bf16 | seg2-stride4 | 0.874 / 10.16 | 1.532 / 10.12 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 2.782 / 13.75 | 2.722 / 13.70 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 4.262 / 2.98 | 4.393 / 2.94 | +| int8-sdnq-hadamard | interval2 | 3.547 / 8.26 | 3.380 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4 | 3.366 / 8.26 | 3.457 / 8.21 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 2.432 / 14.03 | 2.303 / 13.98 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.920 / 3.26 | 4.008 / 3.22 | +| fp8-torchao | interval2 | 3.316 / 8.54 | 2.973 / 8.49 | +| fp8-torchao | seg2-stride4 | 2.994 / 8.54 | 3.046 / 8.49 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HeartMula + +Example: `heartmula.peft-lora`. 音频 token 训练。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### HiDream + +Example: `hidream.peft-lora`. Resolution: 512x512. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.515 / 44.58 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.928 / 33.57 | 1.043 / 33.52 | +| bf16 | interval2 | 0.971 / 33.57 | 1.002 / 33.52 | +| bf16 | seg2-stride4 | 0.952 / 33.57 | 1.010 / 33.52 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | not measured | not measured | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.497 / 17.99 | 2.058 / 17.93 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | not measured | not measured | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 4.920 / 18.10 | 4.338 / 18.05 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +SDNQ Hadamard numbers use `sdnq_compile_mode=eager`. The compiled SDNQ path quantizes HiDream, but this sweep spent the first training step in Inductor dequantizer compilation, so it is not listed as a throughput row. + +### HunyuanVideo + +Example: `hunyuanvideo-1.5-t2v.peft-lora`. Training shape: 480 pixel-area video buckets, 48 frames, batch 2. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | not run | not run | +| bf16 | layer | 7.682 / 26.35 | 22.816 / 26.30 | +| bf16 | interval2 | 7.398 / 26.11 | 22.772 / 26.06 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 10.679 / 58.37 | not run | +| int8-sdnq-hadamard | none | not run | not run | +| int8-sdnq-hadamard | activation-offload | not run | not run | +| int8-sdnq-hadamard | layer | 11.765 / 25.96 | 34.464 / 25.92 | +| int8-sdnq-hadamard | interval2 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4 | not run | not run | +| int8-sdnq-hadamard | seg2-stride4-offload | not run | not run | +| fp8-torchao | none | not run | not run | +| fp8-torchao | activation-offload | not run | not run | +| fp8-torchao | layer | 10.516 / 33.55 | 32.003 / 33.53 | +| fp8-torchao | interval2 | not run | not run | +| fp8-torchao | seg2-stride4 | not run | not run | +| fp8-torchao | seg2-stride4-offload | not run | not run | + +HunyuanVideo 在这个 training shape 下 activation 压力很高。Per-block 和 interval-2 checkpointing 都能稳定运行;不开 checkpointing 时 80 GB H100 也放不下。`seg2-stride4` 在本轮 sweep 中只有启用 attention activation offload 才能跑通,这一行更适合作为内存 fallback,而不是速度推荐。SDNQ Hadamard 可以工作,但可变 conditioning shape 仍会在测量窗口内触发 dynamic kernel compilation。 + +### Ideogram 4.0 + +示例:`ideogram-fp8.peft-lora`。分辨率:1024x1024。`fp8` flavour 使用 Ideogram 4 的原生 weight-only fp8 checkpoint(`base_model_precision=no_change`)。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| fp8-native | none | failed | OOM | +| fp8-native | activation-offload | unsupported | unsupported | +| fp8-native | layer | 1.033 / 12.57 | 3.101 / 11.82 | +| fp8-native | interval2 | 1.030 / 12.33 | 3.098 / 11.82 | +| fp8-native | seg2-stride4 | 1.031 / 12.33 | 3.088 / 11.82 | +| fp8-native | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.702 / 61.78 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.033 / 12.33 | 3.030 / 11.82 | +| int8-sdnq-hadamard | interval2 | 1.032 / 12.33 | 3.034 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4 | 1.028 / 12.33 | 3.035 / 11.82 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | + +### Kandinsky 5 Image + +Example: `kandinsky5-image-6b-t2i.lycoris-lokr`. Resolution: 1024x1024. + +Note: batch 3 at 1024x1024 needs full checkpointing on both cards. SDNQ with Hadamard is the best low-VRAM row; H100 can also use partial checkpointing with SDNQ, but only near the top of the card. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.956 / 25.52 | 10.189 / 25.42 | +| bf16 | layer | 6.590 / 25.58 | 9.458 / 25.55 | +| bf16 | interval2 | OOM | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 7.186 / 20.00 | 10.141 / 19.95 | +| int8-sdnq-hadamard | layer | 6.830 / 20.12 | 9.362 / 20.08 | +| int8-sdnq-hadamard | interval2 | 5.746 / 75.67 | OOM | +| int8-sdnq-hadamard | seg2-stride4 | OOM | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 6.057 / 71.46 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 8.319 / 24.29 | 14.716 / 24.24 | +| fp8-torchao | layer | 7.976 / 24.40 | 13.949 / 24.35 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kandinsky 5 Video + +Example: `kandinsky5-video-2b-t2v.peft-lora`. Resolution: 768x512, 81f. + +Kandinsky 5 video is activation-heavy at this frame count. Full block checkpointing is the practical baseline on both cards. On H100, `interval2` and `seg2-stride4` are faster when they fit; on L40S, SDNQ `interval2` is the only partial-checkpoint row here that fits without attention activation offload. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 2.580 / 9.49 | 7.136 / 9.45 | +| bf16 | layer | 2.267 / 9.83 | 6.379 / 9.79 | +| bf16 | interval2 | 1.967 / 44.57 | OOM | +| bf16 | seg2-stride4 | 1.971 / 46.62 | OOM | +| bf16 | seg2-stride4-offload | 2.275 / 37.99 | 6.249 / 37.94 | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 2.844 / 8.07 | 7.234 / 8.02 | +| int8-sdnq-hadamard | layer | 2.460 / 8.40 | 6.509 / 8.35 | +| int8-sdnq-hadamard | interval2 | 2.126 / 43.12 | 5.641 / 43.08 | +| int8-sdnq-hadamard | seg2-stride4 | 2.125 / 45.19 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 2.451 / 36.56 | 6.322 / 36.51 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 3.867 / 15.11 | 11.501 / 15.06 | +| fp8-torchao | layer | 3.579 / 15.28 | 10.822 / 15.24 | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | OOM | OOM | + +### Kolors + +Example: `kolors.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.635 / 7.22 | 0.628 / 7.17 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.118 / 5.36 | 1.065 / 5.31 | +| bf16 | interval2 | 1.105 / 5.36 | 1.068 / 5.31 | +| bf16 | seg2-stride4 | 1.110 / 5.36 | 1.072 / 5.31 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.860 / 5.68 | 1.726 / 5.63 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.967 / 3.38 | 2.803 / 3.33 | +| int8-sdnq-hadamard | interval2 | 2.770 / 3.32 | 2.655 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4 | 2.805 / 3.32 | 2.742 / 3.27 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.629 / 8.79 | 1.637 / 8.75 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.044 / 3.50 | 3.013 / 3.45 | +| fp8-torchao | interval2 | 3.040 / 3.50 | 3.019 / 3.45 | +| fp8-torchao | seg2-stride4 | 2.988 / 3.50 | 2.935 / 3.45 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Krea 2 + +示例:`krea2.peft-lora`。训练分辨率:512 px 方形裁剪。示例验证设置:1024x1024;benchmark 中已禁用 validation。 + +主表使用 regional compilation;它对 Krea2 step speed 有帮助,但不是干净的 VRAM 对比。compiled graph/workspace 会让几个 checkpoint 模式的峰值接近未 checkpoint 的峰值。关闭 regional compilation 的 bf16 control run 显示 checkpointing 已接通,并呈现预期的 memory/speed 形态: + +这里的 `activation-offload` 行表示 full-block checkpointing 加 attention activation offload。和单独的 full-block `layer` checkpointing 相比,attention offload 在这个 Krea2 shape 下没有降低 peak VRAM,主要增加了 CPU transfer overhead。 + +| Mode | H100 no-compile | L40S no-compile | +| --- | ---: | ---: | +| none | 0.272 / 40.09 | 0.661 / 40.01 | +| layer | 0.371 / 30.06 | 0.919 / 30.01 | +| seg2-stride4 | 0.317 / 34.75 | 0.788 / 34.70 | +| activation-offload | 0.657 / 30.50 | 1.341 / 30.30 | + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.275 / 40.06 | 0.662 / 40.01 | +| bf16 | activation-offload | 0.416 / 34.62 | 0.822 / 34.57 | +| bf16 | layer | 0.268 / 40.06 | 0.665 / 40.01 | +| bf16 | interval2 | 0.274 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4 | 0.279 / 40.06 | 0.663 / 40.01 | +| bf16 | seg2-stride4-offload | 0.404 / 34.62 | 0.819 / 34.57 | +| int8-sdnq-hadamard | none | 0.462 / 27.01 | 0.807 / 26.96 | +| int8-sdnq-hadamard | activation-offload | 0.773 / 21.57 | 1.012 / 21.53 | +| int8-sdnq-hadamard | layer | 0.473 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | interval2 | 0.474 / 27.01 | 0.804 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4 | 0.472 / 27.01 | 0.802 / 26.96 | +| int8-sdnq-hadamard | seg2-stride4-offload | 0.744 / 21.57 | 1.007 / 21.53 | +| fp8-torchao | none | 0.689 / 51.63 | OOM | +| fp8-torchao | activation-offload | 0.975 / 37.77 | 2.058 / 37.73 | +| fp8-torchao | layer | 0.689 / 51.63 | OOM | +| fp8-torchao | interval2 | 0.684 / 51.63 | OOM | +| fp8-torchao | seg2-stride4 | 0.674 / 51.63 | OOM | +| fp8-torchao | seg2-stride4-offload | 0.965 / 37.77 | 2.053 / 37.73 | + +### LongCat Image + +示例:`longcat-image.peft-lora`。训练分辨率:512 px 正方形;验证分辨率:1024x1024。各行使用 `attention_mechanism=native-flash`。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.193 / 16.47 | 0.262 / 16.42 | +| bf16 | activation-offload | 0.544 / 12.73 | 0.543 / 12.69 | +| bf16 | layer | 0.327 / 12.38 | 0.370 / 12.34 | +| bf16 | interval2 | 0.257 / 14.38 | 0.313 / 14.34 | +| bf16 | seg2-stride4 | 0.263 / 14.36 | 0.316 / 14.31 | +| bf16 | seg2-stride4-offload | 0.446 / 13.42 | 0.492 / 13.38 | +| int8-sdnq-hadamard | none | 0.578 / 12.54 | 0.537 / 12.45 | +| int8-sdnq-hadamard | activation-offload | 1.184 / 7.43 | 1.185 / 7.39 | +| int8-sdnq-hadamard | layer | 0.901 / 7.19 | 0.911 / 7.09 | +| int8-sdnq-hadamard | interval2 | 0.718 / 9.73 | 0.695 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4 | 0.735 / 9.72 | 0.717 / 9.68 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.056 / 8.49 | 1.011 / 8.45 | +| fp8-torchao | none | 0.602 / 25.19 | 0.834 / 25.14 | +| fp8-torchao | activation-offload | 1.662 / 7.75 | 1.844 / 7.70 | +| fp8-torchao | layer | 0.984 / 7.57 | 1.080 / 7.53 | +| fp8-torchao | interval2 | 0.750 / 16.15 | 0.938 / 16.10 | +| fp8-torchao | seg2-stride4 | 0.760 / 16.13 | 0.961 / 16.09 | +| fp8-torchao | seg2-stride4-offload | 1.287 / 13.02 | 1.653 / 12.98 | + +### LongCat Video + +示例:`longcat-video.peft-lora+ramtorch`。分辨率:832x480,81f。各行使用 `attention_mechanism=native-flash`。 + +LongCat Video 在这个 shape 下 activation 压力很高。Full per-block checkpointing 是实际可用的基线。Partial checkpoint 行(`interval2`、`seg2-stride4`)在这里放不下;即使 strided 行启用 attention activation offload 也不够。单独的 attention activation offload 可让 bf16 和 SDNQ 跑通,但明显慢于 full checkpointing。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM (77.86 GiB) | OOM (43.30 GiB) | +| bf16 | activation-offload | 25.774 / 37.36 | 49.149 / 37.14 | +| bf16 | layer | 7.448 / 23.73 | 24.866 / 23.68 | +| bf16 | interval2 | OOM (76.41 GiB) | OOM (42.80 GiB) | +| bf16 | seg2-stride4 | OOM (76.42 GiB) | OOM (43.06 GiB) | +| bf16 | seg2-stride4-offload | OOM (76.72 GiB) | OOM (42.29 GiB) | +| int8-sdnq-hadamard | none | OOM (77.40 GiB) | OOM (43.59 GiB) | +| int8-sdnq-hadamard | activation-offload | 30.887 / 35.28 | 61.270 / 35.24 | +| int8-sdnq-hadamard | layer | 8.444 / 21.60 | 25.164 / 21.55 | +| int8-sdnq-hadamard | interval2 | OOM (76.01 GiB) | OOM (42.47 GiB) | +| int8-sdnq-hadamard | seg2-stride4 | OOM (77.01 GiB) | OOM (43.02 GiB) | +| int8-sdnq-hadamard | seg2-stride4-offload | OOM (76.42 GiB) | OOM (42.57 GiB) | +| fp8-torchao | none | OOM (75.87 GiB) | OOM (41.28 GiB) | +| fp8-torchao | activation-offload | 30.163 / 47.88 | OOM (40.11 GiB) | +| fp8-torchao | layer | 8.343 / 34.16 | 24.659 / 34.07 | +| fp8-torchao | interval2 | OOM (74.63 GiB) | OOM (40.85 GiB) | +| fp8-torchao | seg2-stride4 | OOM (75.37 GiB) | OOM (41.19 GiB) | +| fp8-torchao | seg2-stride4-offload | OOM (75.55 GiB) | OOM (40.94 GiB) | + +### LTXVideo 0.9.5 + +Example: `ltxvideo-0.9.5-t2v.peft-lora`. Resolution: 768x512, 49f. + +数字是 warm seconds/step / peak GiB。完整运行平均值包含 setup 和 compile overhead,保存在 sweep artifacts 里。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.274 / 8.77 | 0.275 / 8.72 | +| bf16 | layer | 0.462 / 4.21 | 0.459 / 4.16 | +| bf16 | interval2 | 0.449 / 4.30 | 0.433 / 4.25 | +| bf16 | seg2-stride4 | 0.359 / 6.40 | 0.357 / 6.35 | +| int8-sdnq-hadamard | none | 0.688 / 7.05 | 0.640 / 6.90 | +| int8-sdnq-hadamard | layer | 1.094 / 2.59 | 1.081 / 2.44 | +| int8-sdnq-hadamard | interval2 | 1.112 / 2.57 | 1.073 / 2.53 | +| int8-sdnq-hadamard | seg2-stride4 | 0.887 / 4.61 | 0.817 / 4.56 | +| fp8-torchao | none | 0.655 / 17.45 | 0.735 / 17.41 | +| fp8-torchao | layer | 1.226 / 2.94 | 1.206 / 2.89 | +| fp8-torchao | interval2 | 1.259 / 3.40 | 1.233 / 3.35 | +| fp8-torchao | seg2-stride4 | 0.933 / 9.96 | 0.901 / 9.91 | +| fp8wo-torchao | none | 0.328 / 10.06 | 0.325 / 10.01 | +| fp8wo-torchao | layer | 0.567 / 2.64 | 0.540 / 2.59 | +| fp8wo-torchao | interval2 | 0.555 / 2.84 | 0.531 / 2.79 | +| fp8wo-torchao | seg2-stride4 | 0.443 / 6.20 | 0.432 / 6.15 | + +这个 sweep 中,LTXVideo 0.9 不支持 attention activation offload rows。 + +### LTXVideo2 2.3 + +Example: `ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard`. Resolution: 1280x704, 49f. + +Note: LTXVideo2 2.3 should be read from the no-regional-compile rows in this sweep; regional compile raised memory pressure for this model. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 6.350 / 58.36 | OOM | +| bf16 | layer | 3.993 / 47.95 | OOM | +| bf16 | interval2 | 3.977 / 48.83 | OOM | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | 5.554 / 75.64 | OOM | +| int8-sdnq-hadamard | none | OOM | OOM | +| int8-sdnq-hadamard | activation-offload | 10.102 / 38.20 | 9.852 / 38.15 | +| int8-sdnq-hadamard | layer | 7.753 / 27.78 | 7.579 / 27.73 | +| int8-sdnq-hadamard | interval2 | 7.733 / 28.66 | 7.288 / 28.61 | +| int8-sdnq-hadamard | seg2-stride4 | 6.522 / 61.50 | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | 8.659 / 55.48 | OOM | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | failed | 22.917 / 38.57 | +| fp8-torchao | layer | 8.580 / 30.27 | 10.381 / 30.22 | +| fp8-torchao | interval2 | 8.660 / 33.71 | 10.661 / 33.66 | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | failed | OOM | + +### Lumina2 + +Example: `lumina2.peft-lora`. Resolution: 512x512. + +说明:Lumina2 现在使用 segmented whole-block 路径。`interval2` 会 checkpoint 每个 two-block segment;`seg2-stride4` 会 checkpoint 两个 blocks,让接下来的两个 blocks 保留 activations,然后重复。Attention activation offload 不在这次 Lumina2 run 中。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.235 / 15.99 | 0.332 / 15.94 | +| bf16 | layer | 0.384 / 6.60 | 0.457 / 6.56 | +| bf16 | interval2 | 0.356 / 6.87 | 0.424 / 6.82 | +| bf16 | seg2-stride4 | 0.295 / 11.43 | 0.377 / 11.38 | +| int8-sdnq-hadamard | none | 0.584 / 13.59 | 0.541 / 13.55 | +| int8-sdnq-hadamard | layer | 0.899 / 4.21 | 0.827 / 4.16 | +| int8-sdnq-hadamard | interval2 | 0.865 / 4.48 | 0.835 / 4.43 | +| int8-sdnq-hadamard | seg2-stride4 | 0.719 / 9.03 | 0.707 / 8.99 | +| fp8-torchao | none | 0.598 / 28.03 | 0.763 / 27.98 | +| fp8-torchao | layer | 0.950 / 5.93 | 0.967 / 5.88 | +| fp8-torchao | interval2 | 0.938 / 6.70 | 0.974 / 6.66 | +| fp8-torchao | seg2-stride4 | 0.765 / 17.36 | 0.901 / 17.32 | +| fp8wo-torchao | none | 0.273 / 17.97 | 0.389 / 17.93 | +| fp8wo-torchao | layer | 0.452 / 4.99 | 0.525 / 4.94 | +| fp8wo-torchao | interval2 | 0.427 / 5.40 | 0.522 / 5.35 | +| fp8wo-torchao | seg2-stride4 | 0.360 / 11.68 | 0.466 / 11.64 | + +### MageFlow + +Example: `mageflow-image-24g.peft-lora`. Resolution: 1024x1024. + +说明:MageFlow 的 1024px 可变形状图像路径主要受益于 attention activation offload 和 weight-only FP8。Block checkpointing 模式是可用的,但在本轮 sweep 中没有降低测得的峰值驻留显存。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 7.735 / 36.74 | 7.658 / 36.69 | +| bf16 | activation-offload | 8.902 / 23.05 | 9.477 / 23.00 | +| bf16 | layer | 7.805 / 36.74 | 7.504 / 36.69 | +| bf16 | interval2 | 8.036 / 36.74 | 7.762 / 36.69 | +| bf16 | seg2-stride4 | 7.882 / 36.74 | 7.531 / 36.69 | +| bf16 | seg2-stride4-offload | 8.833 / 23.05 | 9.288 / 23.01 | +| int8-sdnq-hadamard | none | 81.016 / 37.38 | 94.991 / 37.34 | +| fp8wo-torchao | none | 5.454 / 36.86 | 5.772 / 36.82 | +| fp8wo-torchao | activation-offload | 6.295 / 23.18 | 6.738 / 23.14 | +| fp8wo-torchao | seg2-stride4 | 5.542 / 36.86 | 5.595 / 36.82 | + +### OmniGen + +Example: `omnigen.lycoris-lokr`. Resolution: 1024x1024. + +说明:OmniGen 使用 token-ID prompt,而不是缓存的文本 embedding。这里测的是已支持的无 checkpoint 和整 block torch checkpointing 路径;interval、segmented-stride 和 attention-offload 控制在本次 sweep 中尚未为这个 family 实现。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.425 / 14.20 | 0.293 / 14.24 | +| bf16 | layer | 0.597 / 10.13 | 0.389 / 10.09 | +| int8-sdnq-hadamard | none | 1.523 / 11.00 | 1.388 / 11.06 | +| int8-sdnq-hadamard | layer | 1.312 / 6.75 | 1.064 / 6.70 | +| fp8-torchao | none | 0.690 / 19.25 | 0.608 / 19.30 | +| fp8-torchao | layer | 1.069 / 7.09 | 0.824 / 7.04 | +| fp8wo-torchao | none | 0.454 / 17.71 | 0.377 / 17.73 | +| fp8wo-torchao | layer | 0.646 / 6.98 | 0.534 / 6.94 | + +### PixArt + +Example: `pixart.lycoris-lokr`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.700 / 41.73 | 1.734 / 41.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 2.433 / 5.63 | 2.346 / 5.58 | +| bf16 | interval2 | 2.440 / 6.01 | 2.348 / 5.96 | +| bf16 | seg2-stride4 | 2.092 / 23.90 | 2.072 / 23.86 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.905 / 47.58 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.738 / 6.07 | 2.937 / 6.03 | +| int8-sdnq-hadamard | interval2 | 2.734 / 6.03 | 2.943 / 5.99 | +| int8-sdnq-hadamard | seg2-stride4 | 2.336 / 26.85 | 2.596 / 26.81 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.254 / 8.06 | 4.646 / 8.01 | +| fp8-torchao | interval2 | 3.260 / 9.66 | 4.649 / 9.61 | +| fp8-torchao | seg2-stride4 | 2.827 / 63.27 | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Qwen Image + +Example: `qwen_image.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.202 / 41.04 | 3.377 / 40.99 | +| bf16 | interval2 | 1.201 / 41.04 | 3.382 / 40.99 | +| bf16 | seg2-stride4 | 1.205 / 41.03 | 3.385 / 40.99 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.675 / 63.48 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.640 / 24.09 | 3.928 / 24.05 | +| int8-sdnq-hadamard | interval2 | 2.722 / 24.09 | 3.918 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4 | 2.663 / 24.09 | 3.919 / 24.05 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 3.088 / 25.34 | 6.172 / 25.29 | +| fp8-torchao | interval2 | 3.095 / 25.34 | 6.173 / 25.29 | +| fp8-torchao | seg2-stride4 | 3.125 / 25.34 | 6.141 / 25.29 | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Sana + +Example: `sana.lycoris-lokr`. Resolution: 1024x1024. + +Note: Sana has interval checkpointing; stride is not a separate segmented schedule for this family in the measured rows. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 23.71 | 0.597 / 23.67 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.529 / 23.72 | 0.596 / 23.66 | +| bf16 | interval2 | 0.530 / 23.72 | 0.598 / 23.66 | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.554 / 22.92 | 0.590 / 22.88 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.554 / 22.92 | 0.589 / 22.88 | +| int8-sdnq-hadamard | interval2 | 0.556 / 22.92 | 0.591 / 22.88 | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.633 / 33.67 | 0.753 / 33.62 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.633 / 33.67 | 0.755 / 33.62 | +| fp8-torchao | interval2 | 0.631 / 33.67 | 0.759 / 33.62 | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SanaVideo + +Example: `sanavideo-2b-480p.peft-lora`. Resolution: 832x480, 49f. + +Note: SanaVideo 使用 linear attention,所以 attention activation offload 仍不支持。标准路径支持 segmented whole-block checkpointing。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.599 / 59.15 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.597 / 59.15 | OOM | +| bf16 | interval2 | not run | OOM | +| bf16 | seg2-stride4 | not run | OOM | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.641 / 58.36 | OOM | +| int8-sdnq-hadamard | interval2 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4 | not run | OOM | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | OOM | OOM | +| fp8-torchao | interval2 | OOM | OOM | +| fp8-torchao | seg2-stride4 | OOM | OOM | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD 1.x + +Example: `sd1x-dreamshaper.peft-lora`. Resolution: 512x512. + +说明:SD1x 使用 diffusers UNet path。普通 layer checkpointing 受支持,但 interval 和 segmented stride controls 没有接入这个 family。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.181 / 2.87 | 0.176 / 2.83 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.305 / 1.98 | 0.292 / 1.93 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.446 / 3.07 | 0.431 / 3.04 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 0.701 / 1.79 | 0.666 / 1.74 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 0.410 / 4.23 | 0.401 / 4.18 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 0.785 / 1.96 | 0.756 / 1.91 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### SD3 + +Example: `sd3.peft-lora`. Resolution: 1024x1024. + +说明:SD3 在普通 transformer path 上使用真正的连续 segmented checkpointing。Attention activation offload 已支持;它会大幅降低 VRAM,但会牺牲 throughput。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.529 / 34.59 | 1.189 / 34.53 | +| bf16 | activation-offload | 1.335 / 9.68 | 2.850 / 9.63 | +| bf16 | layer | 0.721 / 7.67 | 1.607 / 7.62 | +| bf16 | interval2 | 0.723 / 8.93 | 1.606 / 8.88 | +| bf16 | seg2-stride4 | 0.620 / 20.14 | 1.398 / 20.09 | +| bf16 | seg2-stride4-offload | 1.229 / 12.21 | 2.639 / 12.16 | +| int8-sdnq-hadamard | none | 0.858 / 33.66 | 1.381 / 33.61 | +| int8-sdnq-hadamard | activation-offload | 1.845 / 8.90 | 3.297 / 8.85 | +| int8-sdnq-hadamard | layer | 1.265 / 6.58 | 1.857 / 6.53 | +| int8-sdnq-hadamard | interval2 | 1.264 / 7.30 | 1.858 / 7.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.049 / 19.02 | 1.625 / 18.98 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.527 / 12.42 | 3.023 / 12.37 | +| fp8-torchao | none | OOM | OOM | +| fp8-torchao | activation-offload | 4.707 / 10.27 | 9.814 / 10.23 | +| fp8-torchao | layer | 1.575 / 9.13 | 3.279 / 9.09 | +| fp8-torchao | interval2 | 1.576 / 12.62 | 3.282 / 12.58 | +| fp8-torchao | seg2-stride4 | 1.376 / 45.70 | OOM | +| fp8-torchao | seg2-stride4-offload | 4.184 / 26.17 | 8.714 / 26.12 | +| fp8wo-torchao | none | 0.567 / 35.21 | 1.252 / 35.17 | +| fp8wo-torchao | activation-offload | 1.392 / 8.71 | 2.921 / 8.67 | +| fp8wo-torchao | layer | 0.795 / 5.69 | 1.736 / 5.65 | +| fp8wo-torchao | interval2 | 0.797 / 7.08 | 1.734 / 7.03 | +| fp8wo-torchao | seg2-stride4 | 0.679 / 19.43 | 1.490 / 19.38 | +| fp8wo-torchao | seg2-stride4-offload | 1.257 / 12.00 | 2.676 / 11.96 | + +### SDXL + +Example: `sdxl.lycoris-lokr`. Resolution: 1024x1024. + +Note: SDXL has real layer checkpointing. Interval and stride rows are included as coverage data, not as segmented-support recommendations. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.606 / 13.03 | 0.585 / 12.98 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.080 / 6.53 | 1.029 / 6.48 | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 1.741 / 13.72 | 1.643 / 13.68 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 2.820 / 4.64 | 2.647 / 4.59 | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | 1.608 / 26.22 | 1.582 / 26.16 | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | 2.939 / 5.10 | 2.890 / 5.04 | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Stable Cascade + +Example: `cascade-stage-c.lycoris-lokr`. Resolution: 1024x1024. + +说明:stage C 使用全精度 prior 路径。这些行使用 `mixed_precision=no` 和 `base_model_precision=no_change`;量化 base precision 行对这个模型没有实际意义。interval 和 stride 模式作用在 UNet 的 Res/Timestep/Attention micro-block 序列上。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.884 / 51.52 | OOM | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 1.179 / 22.68 | 2.135 / 22.61 | +| bf16 | interval2 | 1.032 / 36.99 | 1.871 / 36.92 | +| bf16 | seg2-stride4 | 1.032 / 37.20 | 1.870 / 37.13 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | unsupported | unsupported | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | unsupported | unsupported | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Wan 2.1 T2V 1.3B + +Example: `wan2.1-t2v-1.3b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 1.3B should be read from the no-regional-compile/RamTorch rows; regional compile was not a useful throughput setting in this sweep. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 1.407 / 71.90 | OOM | +| bf16 | activation-offload | 3.472 / 8.78 | 7.179 / 8.66 | +| bf16 | layer | 2.099 / 4.73 | 4.459 / 4.68 | +| bf16 | interval2 | 2.139 / 6.32 | 4.514 / 6.27 | +| bf16 | seg2-stride4 | 1.806 / 39.25 | 3.921 / 39.21 | +| bf16 | seg2-stride4-offload | 2.993 / 22.66 | 6.493 / 22.61 | +| int8-sdnq-hadamard | none | 1.850 / 71.91 | OOM | +| int8-sdnq-hadamard | activation-offload | 4.204 / 8.72 | 7.387 / 8.68 | +| int8-sdnq-hadamard | layer | 2.790 / 4.70 | 4.989 / 4.65 | +| int8-sdnq-hadamard | interval2 | 2.874 / 6.29 | 5.093 / 6.24 | +| int8-sdnq-hadamard | seg2-stride4 | 2.558 / 39.27 | 4.393 / 39.22 | +| int8-sdnq-hadamard | seg2-stride4-offload | 3.711 / 22.67 | 6.695 / 22.63 | +| fp8-torchao | none | 1.727 / 73.57 | OOM | +| fp8-torchao | activation-offload | 4.061 / 10.08 | 7.404 / 9.96 | +| fp8-torchao | layer | 2.607 / 5.98 | 4.888 / 5.93 | +| fp8-torchao | interval2 | 2.744 / 7.57 | 4.916 / 7.52 | +| fp8-torchao | seg2-stride4 | 2.246 / 40.55 | 4.245 / 40.50 | +| fp8-torchao | seg2-stride4-offload | 3.602 / 24.02 | 6.683 / 23.91 | + +### Wan 2.1 T2V 14B + +Example: `wan2.1-t2v-14b-480p-single-gpu.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan 14B is mainly a fit test for activation savings. Status-only cells are still useful because they show which combinations reached the memory limit. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | OOM | OOM | +| bf16 | activation-offload | 13.144 / 36.80 | OOM | +| bf16 | layer | 7.162 / 16.28 | 21.770 / 16.23 | +| bf16 | interval2 | 7.172 / 19.62 | 21.777 / 19.58 | +| bf16 | seg2-stride4 | OOM | OOM | +| bf16 | seg2-stride4-offload | OOM | OOM | +| int8-sdnq-hadamard | none | unsupported | unsupported | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | unsupported | unsupported | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | failed | failed | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | failed | failed | +| fp8-torchao | seg2-stride4 | failed | failed | +| fp8-torchao | seg2-stride4-offload | failed | failed | + +### Wan S2V + +Example: `wan-s2v-14b-480p.peft-lora+ramtorch`. Resolution: 832x480, 81f. + +Note: Wan S2V is included as coverage data for the video/audio path. Treat failed cells as implementation coverage gaps. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | failed | failed | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | failed | failed | +| bf16 | interval2 | unsupported | unsupported | +| bf16 | seg2-stride4 | unsupported | unsupported | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | failed | failed | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | failed | failed | +| int8-sdnq-hadamard | interval2 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4 | unsupported | unsupported | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8-torchao | none | failed | failed | +| fp8-torchao | activation-offload | unsupported | unsupported | +| fp8-torchao | layer | failed | failed | +| fp8-torchao | interval2 | unsupported | unsupported | +| fp8-torchao | seg2-stride4 | unsupported | unsupported | +| fp8-torchao | seg2-stride4-offload | unsupported | unsupported | + +### Z-Image Turbo + +Example: `z-image-turbo.peft-lora`. Resolution: 1024x1024. + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.243 / 21.25 | 0.316 / 21.21 | +| bf16 | activation-offload | 0.837 / 13.24 | 0.805 / 13.19 | +| bf16 | layer | 0.479 / 12.87 | 0.493 / 12.83 | +| bf16 | interval2 | 0.452 / 13.04 | 0.477 / 12.99 | +| bf16 | seg2-stride4 | 0.349 / 16.88 | 0.400 / 16.83 | +| bf16 | seg2-stride4-offload | 0.681 / 15.03 | 0.736 / 14.99 | +| int8-sdnq-hadamard | none | 0.645 / 15.60 | 0.615 / 15.55 | +| int8-sdnq-hadamard | activation-offload | 1.439 / 7.61 | 1.382 / 7.56 | +| int8-sdnq-hadamard | layer | 1.046 / 7.25 | 1.021 / 7.20 | +| int8-sdnq-hadamard | interval2 | 1.074 / 7.41 | 0.996 / 7.36 | +| int8-sdnq-hadamard | seg2-stride4 | 0.867 / 11.25 | 0.841 / 11.20 | +| int8-sdnq-hadamard | seg2-stride4-offload | 1.202 / 9.39 | 1.162 / 9.35 | +| fp8-torchao | none | 1.232 / 37.50 | 1.476 / 37.46 | +| fp8-torchao | activation-offload | 3.623 / 7.97 | 3.564 / 7.93 | +| fp8-torchao | layer | 2.319 / 7.93 | 2.344 / 7.88 | +| fp8-torchao | interval2 | 2.336 / 8.80 | 2.309 / 8.75 | +| fp8-torchao | seg2-stride4 | 1.843 / 22.55 | 1.930 / 22.50 | +| fp8-torchao | seg2-stride4-offload | 2.947 / 15.57 | 3.243 / 15.52 | + +### ZLab I1 + +Example: `zlab-i1.peft-lora`. Resolution: 1024x1024. + +说明:ZLab I1 会把 U-Net 式 skip tensors 放进 segmented checkpoint state。该 family 尚未接入 attention activation offload。 + +| Precision | Mode | H100 | L40S | +| --- | --- | ---: | ---: | +| bf16 | none | 0.462 / 22.21 | 0.865 / 22.16 | +| bf16 | activation-offload | unsupported | unsupported | +| bf16 | layer | 0.693 / 7.79 | 1.148 / 7.75 | +| bf16 | interval2 | 0.676 / 8.30 | 1.152 / 8.25 | +| bf16 | seg2-stride4 | 0.567 / 14.97 | 1.014 / 14.92 | +| bf16 | seg2-stride4-offload | unsupported | unsupported | +| int8-sdnq-hadamard | none | 0.861 / 19.21 | 0.926 / 19.16 | +| int8-sdnq-hadamard | activation-offload | unsupported | unsupported | +| int8-sdnq-hadamard | layer | 1.385 / 4.78 | 1.265 / 4.74 | +| int8-sdnq-hadamard | interval2 | 1.298 / 5.30 | 1.277 / 5.26 | +| int8-sdnq-hadamard | seg2-stride4 | 1.073 / 11.98 | 1.098 / 11.93 | +| int8-sdnq-hadamard | seg2-stride4-offload | unsupported | unsupported | +| fp8wo-torchao | none | 0.504 / 25.08 | 0.930 / 25.02 | +| fp8wo-torchao | activation-offload | unsupported | unsupported | +| fp8wo-torchao | layer | 0.772 / 5.12 | 1.280 / 5.07 | +| fp8wo-torchao | interval2 | 0.759 / 5.84 | 1.290 / 5.79 | +| fp8wo-torchao | seg2-stride4 | 0.633 / 15.12 | 1.115 / 15.08 | +| fp8wo-torchao | seg2-stride4-offload | unsupported | unsupported | + + diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.es.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.es.md index 3fa6a5619..99f7564c0 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.es.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.es.md @@ -32,9 +32,15 @@ En familias compatibles también puedes checkpointar menos bloques: } ``` -`gradient_checkpointing_interval: 2` checkpointa cada dos bloques compatibles. Valores más altos hacen menos checkpointing y dejan más activaciones en VRAM. +`gradient_checkpointing_interval: 2` checkpointa chunks contiguos de dos bloques en rutas de bloque completo compatibles. Valores más altos recomputan menos y dejan más activaciones en VRAM. -`torch-ffn` y `unsloth-ffn` soportan actualmente bloques estilo Flux.1 y MageFlow. Otras familias fallan claramente hasta que sus bloques expongan el mismo límite seguro. +En esas rutas segmentadas, `gradient_checkpointing_segment_stride` también funciona con `unsloth`. Trátalo como una palanca para caber, no para acelerar: los bloques saltados se quedan en GPU, mientras los bloques checkpointed siguen usando CPU offload para tensores guardados. Para el resumen con torch y benchmarks por modelo, consulta [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md). + +`gradient_checkpointing_offload_attention` es independiente del backend. En bloques compatibles con separación attention/FFN, descarga las activaciones guardadas del lado attention. Puede ejecutarse solo o combinarse con `torch`, `torch-ffn`, `unsloth` o `unsloth-ffn` cuando el modelo soporte ese backend. + +`gradient_checkpointing_offload_pin_memory_max_buckets` controla el pooling de CPU pinned para tensores guardados offloaded. El valor predeterminado es `12` buckets de tensor distintos; usa `0` para usar solo memoria CPU normal. + +`torch-ffn` y `unsloth-ffn` soportan actualmente Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan y Z-Image. Otras familias fallan claramente hasta que sus bloques expongan el mismo límite seguro. ## Qué intercambia diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.hi.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.hi.md index ee3860da5..07007efe1 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.hi.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.hi.md @@ -32,9 +32,15 @@ Supported model families में आप कम blocks भी checkpoint कर } ``` -`gradient_checkpointing_interval: 2` हर दूसरे supported block को checkpoint करता है। Higher values कम checkpointing करती हैं और VRAM में ज्यादा activations रखती हैं। +`gradient_checkpointing_interval: 2` supported whole-block paths पर contiguous दो-block chunks checkpoint करता है। Higher values कम recompute करती हैं और VRAM में ज्यादा activations रखती हैं। -`torch-ffn` और `unsloth-ffn` अभी Flux.1-style blocks और MageFlow support करते हैं। बाकी model families साफ error देंगी जब तक उनके blocks वही safe boundary expose नहीं करते। +इन segmented paths पर `gradient_checkpointing_segment_stride` भी `unsloth` के साथ काम करता है। इसे fit lever मानें, speed lever नहीं: skipped blocks GPU पर रहते हैं, और checkpointed blocks saved tensors के लिए CPU offload use करते हैं। Torch-only overview और model benchmarks के लिए [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md) देखें। + +`gradient_checkpointing_offload_attention` backend से अलग option है। Supported attention/FFN split blocks पर यह attention-side saved activations offload करता है। यह अकेले चल सकता है या model उस backend को support करे तो `torch`, `torch-ffn`, `unsloth`, या `unsloth-ffn` के साथ combine हो सकता है। + +`gradient_checkpointing_offload_pin_memory_max_buckets` offloaded saved tensors के लिए pinned CPU pooling control करता है। default `12` distinct tensor buckets है; normal CPU memory ही इस्तेमाल करने के लिए इसे `0` करें। + +`torch-ffn` और `unsloth-ffn` अभी Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan, और Z-Image support करते हैं। बाकी model families साफ error देंगी जब तक उनके blocks वही safe boundary expose नहीं करते। ## Tradeoff diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.ja.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.ja.md index 24bcba5f7..c44502197 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.ja.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.ja.md @@ -32,9 +32,15 @@ } ``` -`gradient_checkpointing_interval: 2` は対応 block を 1 つおきに checkpoint します。値を大きくすると checkpointing は減り、VRAM に残る activation は増えます。 +`gradient_checkpointing_interval: 2` は対応する whole-block path で連続した 2-block chunk を checkpoint します。値を大きくすると再計算は減り、VRAM に残る activation は増えます。 -`torch-ffn` と `unsloth-ffn` は現在 Flux.1-style blocks と MageFlow に対応しています。他のモデル族は、同じ安全な境界を expose するまで明示的に失敗します。 +これらの segmented path では、`gradient_checkpointing_segment_stride` も `unsloth` で使えます。速度目的ではなく fit lever として扱ってください。Skip された blocks は GPU に残り、checkpoint された blocks は保存 tensor に CPU offload を使います。Torch-only の概要とモデル別 benchmark は [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md) を参照してください。 + +`gradient_checkpointing_offload_attention` は backend とは別の option です。対応する attention/FFN split blocks では、attention 側の保存 activations を offload します。単体でも実行でき、モデルがその backend をサポートする場合は `torch`、`torch-ffn`、`unsloth`、`unsloth-ffn` と組み合わせられます。 + +`gradient_checkpointing_offload_pin_memory_max_buckets` は offload された保存 tensor の pinned CPU pooling を制御します。デフォルトは `12` 個の distinct tensor buckets です。`0` にすると通常の CPU memory だけを使います。 + +`torch-ffn` と `unsloth-ffn` は現在 Chroma、Flux、Krea 2、LTXVideo2、MageFlow、Wan、Z-Image に対応しています。他のモデル族は、同じ安全な境界を expose するまで明示的に失敗します。 ## 何を交換するか diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.md index d60b0556b..170aeb9d1 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.md @@ -32,9 +32,15 @@ For supported model families, you can also checkpoint fewer blocks: } ``` -`gradient_checkpointing_interval: 2` checkpoints every other supported block. Higher values spend less time checkpointing and keep more activations in VRAM. +`gradient_checkpointing_interval: 2` checkpoints contiguous two-block chunks on supported whole-block paths. Higher values spend less time recomputing and keep more activations in VRAM. -`torch-ffn` and `unsloth-ffn` currently support Flux.1-style blocks and MageFlow. Other model families fail clearly until their block internals expose the same safe boundary. +On those segmented paths, `gradient_checkpointing_segment_stride` also works with `unsloth`. Treat it as a fit lever, not a speed lever: the skipped blocks stay on GPU, while checkpointed blocks still use CPU offload for saved tensors. For the torch-only overview and model benchmarks, see [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md). + +`gradient_checkpointing_offload_attention` is separate from the backend. On supported attention/FFN split blocks, it offloads attention-side saved activations. It can run by itself or combine with `torch`, `torch-ffn`, `unsloth`, or `unsloth-ffn` when that backend is supported by the model. + +`gradient_checkpointing_offload_pin_memory_max_buckets` controls pinned CPU pooling for offloaded saved tensors. The default is `12` distinct tensor buckets; set it to `0` to use normal CPU memory only. + +`torch-ffn` and `unsloth-ffn` currently support Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan, and Z-Image. Other model families fail clearly until their block internals expose the same safe boundary. ## What It Trades diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.pt-BR.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.pt-BR.md index 9689da991..f2594e44c 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.pt-BR.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.pt-BR.md @@ -32,9 +32,15 @@ Em famílias compatíveis, você também pode checkpointar menos blocos: } ``` -`gradient_checkpointing_interval: 2` faz checkpoint a cada dois blocos compatíveis. Valores maiores fazem menos checkpointing e mantêm mais activations na VRAM. +`gradient_checkpointing_interval: 2` faz checkpoint de chunks contiguos de dois blocos em caminhos whole-block compatíveis. Valores maiores recomputam menos e mantêm mais activations na VRAM. -`torch-ffn` e `unsloth-ffn` atualmente suportam blocos estilo Flux.1 e MageFlow. Outras famílias falham claramente até seus blocos exporem a mesma fronteira segura. +Nesses caminhos segmentados, `gradient_checkpointing_segment_stride` também funciona com `unsloth`. Trate como alavanca para caber, não para acelerar: os blocos pulados ficam na GPU, enquanto os blocos checkpointed ainda usam CPU offload para tensores salvos. Para o resumo torch-only e benchmarks por modelo, veja [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md). + +`gradient_checkpointing_offload_attention` e independente do backend. Em blocos compatíveis com separacao attention/FFN, faz offload das activations salvas do lado attention. Pode rodar sozinho ou ser combinado com `torch`, `torch-ffn`, `unsloth` ou `unsloth-ffn` quando o modelo suportar esse backend. + +`gradient_checkpointing_offload_pin_memory_max_buckets` controla o pooling de CPU pinned para tensores salvos offloaded. O padrao e `12` buckets de tensor distintos; use `0` para usar apenas memoria CPU normal. + +`torch-ffn` e `unsloth-ffn` atualmente suportam Chroma, Flux, Krea 2, LTXVideo2, MageFlow, Wan e Z-Image. Outras famílias falham claramente até seus blocos exporem a mesma fronteira segura. ## O tradeoff diff --git a/documentation/experimental/UNSLOTH_CHECKPOINTING.zh.md b/documentation/experimental/UNSLOTH_CHECKPOINTING.zh.md index 72f1c9c12..3391eb691 100644 --- a/documentation/experimental/UNSLOTH_CHECKPOINTING.zh.md +++ b/documentation/experimental/UNSLOTH_CHECKPOINTING.zh.md @@ -32,9 +32,15 @@ } ``` -`gradient_checkpointing_interval: 2` 表示每两个支持的 block 做一次 checkpoint。值越大,checkpoint 越少,VRAM 里保留的 activation 越多。 +`gradient_checkpointing_interval: 2` 会在受支持的 whole-block 路径上 checkpoint 连续的两个 block chunk。值越大,重算越少,VRAM 里保留的 activation 越多。 -`torch-ffn` 和 `unsloth-ffn` 目前支持 Flux.1 风格 blocks 和 MageFlow。其他模型族会明确报错,直到它们的 block 暴露同样安全的边界。 +在这些分段路径上,`gradient_checkpointing_segment_stride` 也可以和 `unsloth` 一起使用。把它当作 fit lever,而不是 speed lever:跳过的 blocks 仍留在 GPU,checkpointed blocks 仍会把保存 tensor CPU offload。Torch-only 概览和模型 benchmark 见 [Segmented Checkpointing](SEGMENTED_CHECKPOINTING.md)。 + +`gradient_checkpointing_offload_attention` 独立于 backend。在支持 attention/FFN 分离的 blocks 上,它会 offload attention 侧保存的 activations。它可以单独运行;当模型支持所选 backend 时,也可与 `torch`、`torch-ffn`、`unsloth` 或 `unsloth-ffn` 组合。 + +`gradient_checkpointing_offload_pin_memory_max_buckets` 控制 offloaded saved tensors 的 pinned CPU pooling。默认是 `12` 个不同 tensor buckets;设为 `0` 时只使用普通 CPU memory。 + +`torch-ffn` 和 `unsloth-ffn` 目前支持 Chroma、Flux、Krea 2、LTXVideo2、MageFlow、Wan 和 Z-Image。其他模型族会明确报错,直到它们的 block 暴露同样安全的边界。 ## 它交换了什么 diff --git a/documentation/index.es.md b/documentation/index.es.md index 532579eae..546744159 100644 --- a/documentation/index.es.md +++ b/documentation/index.es.md @@ -72,9 +72,9 @@ --- - Funciones de investigación como AnyFlow, cuantización SDNQ Hadamard estilo ConvRot, checkpointing estilo Unsloth, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention y Video CREPA + Funciones de investigación como AnyFlow, cuantización SDNQ Hadamard estilo ConvRot, checkpointing segmentado, checkpointing estilo Unsloth, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention y Video CREPA - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/index.hi.md b/documentation/index.hi.md index c725ff7e7..df6665f2c 100644 --- a/documentation/index.hi.md +++ b/documentation/index.hi.md @@ -72,9 +72,9 @@ --- - AnyFlow, ConvRot-style SDNQ Hadamard quantization, Unsloth-style checkpointing, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention और Video CREPA जैसी research features + AnyFlow, ConvRot-style SDNQ Hadamard quantization, segmented checkpointing, Unsloth-style checkpointing, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention और Video CREPA जैसी research features - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/index.ja.md b/documentation/index.ja.md index 75b99cffb..4d4f3d807 100644 --- a/documentation/index.ja.md +++ b/documentation/index.ja.md @@ -72,9 +72,9 @@ --- - AnyFlow、ConvRot 形式の SDNQ Hadamard 量子化、Unsloth 形式 checkpointing、Prompt2Effect、Self-Flow、Flow-DPO、LayerSync、Diff2Flow、Metal Flash Attention、Video CREPA などの研究向け機能 + AnyFlow、ConvRot 形式の SDNQ Hadamard 量子化、segmented checkpointing、Unsloth 形式 checkpointing、Prompt2Effect、Self-Flow、Flow-DPO、LayerSync、Diff2Flow、Metal Flash Attention、Video CREPA などの研究向け機能 - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/index.md b/documentation/index.md index a45ac1695..c2c2df4da 100644 --- a/documentation/index.md +++ b/documentation/index.md @@ -72,9 +72,9 @@ --- - Research features such as AnyFlow, ConvRot-style SDNQ Hadamard quantization, Unsloth-style checkpointing, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention, and Video CREPA + Research features such as AnyFlow, ConvRot-style SDNQ Hadamard quantization, segmented checkpointing, Unsloth-style checkpointing, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention, and Video CREPA - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/index.pt-BR.md b/documentation/index.pt-BR.md index 222169c21..6b3502dd8 100644 --- a/documentation/index.pt-BR.md +++ b/documentation/index.pt-BR.md @@ -72,9 +72,9 @@ --- - Recursos de pesquisa como AnyFlow, quantização SDNQ Hadamard no estilo ConvRot, checkpointing estilo Unsloth, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention e Video CREPA + Recursos de pesquisa como AnyFlow, quantização SDNQ Hadamard no estilo ConvRot, checkpointing segmentado, checkpointing estilo Unsloth, Prompt2Effect, Self-Flow, Flow-DPO, LayerSync, Diff2Flow, Metal Flash Attention e Video CREPA - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/index.zh.md b/documentation/index.zh.md index 522cff98e..d97af565c 100644 --- a/documentation/index.zh.md +++ b/documentation/index.zh.md @@ -72,9 +72,9 @@ --- - AnyFlow、ConvRot 风格的 SDNQ Hadamard 量化、Unsloth 风格 checkpointing、Prompt2Effect、Self-Flow、Flow-DPO、LayerSync、Diff2Flow、Metal Flash Attention、Video CREPA 等研究功能 + AnyFlow、ConvRot 风格的 SDNQ Hadamard 量化、segmented checkpointing、Unsloth 风格 checkpointing、Prompt2Effect、Self-Flow、Flow-DPO、LayerSync、Diff2Flow、Metal Flash Attention、Video CREPA 等研究功能 - [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) + [:octicons-arrow-right-24: AnyFlow](experimental/ANYFLOW.md) · [:octicons-arrow-right-24: ConvRot / Hadamard SDNQ](experimental/CONVROT.md) · [:octicons-arrow-right-24: Segmented Checkpointing](experimental/SEGMENTED_CHECKPOINTING.md) · [:octicons-arrow-right-24: Unsloth Checkpointing](experimental/UNSLOTH_CHECKPOINTING.md) · [:octicons-arrow-right-24: Metal Flash Attention](experimental/METAL_FLASH_ATTENTION.md) diff --git a/documentation/quickstart/FLUX2.es.md b/documentation/quickstart/FLUX2.es.md index 31f2841bf..dc8b27847 100644 --- a/documentation/quickstart/FLUX2.es.md +++ b/documentation/quickstart/FLUX2.es.md @@ -25,7 +25,7 @@ Para seleccionar una variante, configura `model_flavour` en tu configuración: } ``` -> **Importante**: Para `klein-4b` y `klein-9b`, deja `pretrained_text_encoder_model_name_or_path` sin definir a menos que realmente quieras reemplazar el codificador Qwen3 incluido. Si configuras ese campo, anulas el valor predeterminado de Klein y puedes provocar la descarga de otro codificador de texto. +> **Importante**: Para `klein-4b` y `klein-9b`, deja `qwen_text_encoder_model_name_or_path` sin definir para usar el codificador Qwen3 incluido. Configuralo solo cuando reemplaces el codificador Qwen3 incluido por otra fuente compatible con Qwen. ## Resumen del modelo diff --git a/documentation/quickstart/FLUX2.hi.md b/documentation/quickstart/FLUX2.hi.md index b3c07dc30..5e30432c7 100644 --- a/documentation/quickstart/FLUX2.hi.md +++ b/documentation/quickstart/FLUX2.hi.md @@ -25,7 +25,7 @@ FLUX.2 तीन वेरिएंट में आता है: } ``` -> **महत्वपूर्ण**: `klein-4b` और `klein-9b` के लिए `pretrained_text_encoder_model_name_or_path` को unset छोड़ें, जब तक कि आप bundled Qwen3 text encoder को जानबूझकर बदलना न चाहते हों। इस field को सेट करने पर Klein का default override हो जाता है और किसी दूसरे text encoder का download शुरू हो सकता है। +> **महत्वपूर्ण**: `klein-4b` और `klein-9b` के लिए bundled Qwen3 text encoder उपयोग करने के लिए `qwen_text_encoder_model_name_or_path` को unset छोड़ें। इसे केवल तब सेट करें जब bundled Qwen3 text encoder को किसी दूसरे Qwen-compatible source से बदलना हो। ## मॉडल ओवरव्यू diff --git a/documentation/quickstart/FLUX2.ja.md b/documentation/quickstart/FLUX2.ja.md index 5831ac7c2..5359f2fb3 100644 --- a/documentation/quickstart/FLUX2.ja.md +++ b/documentation/quickstart/FLUX2.ja.md @@ -25,7 +25,7 @@ FLUX.2は3つのバリアントがあります: } ``` -> **重要**: `klein-4b` と `klein-9b` では、同梱のQwen3テキストエンコーダーを意図的に置き換えたい場合を除き、`pretrained_text_encoder_model_name_or_path` は設定しないでください。この項目を設定するとKleinのデフォルトを上書きし、別のテキストエンコーダーのダウンロードが発生することがあります。 +> **重要**: `klein-4b` と `klein-9b` では、同梱の Qwen3 テキストエンコーダーを使用するため `qwen_text_encoder_model_name_or_path` は未設定のままにします。同梱の Qwen3 テキストエンコーダーを別の Qwen 互換ソースに置き換える場合のみ設定してください。 ## モデル概要 diff --git a/documentation/quickstart/FLUX2.md b/documentation/quickstart/FLUX2.md index 5b96d49c5..a71ed561b 100644 --- a/documentation/quickstart/FLUX2.md +++ b/documentation/quickstart/FLUX2.md @@ -25,7 +25,7 @@ To select a variant, set `model_flavour` in your config: } ``` -> **Important**: For `klein-4b` and `klein-9b`, leave `pretrained_text_encoder_model_name_or_path` unset unless you intentionally want to replace the bundled Qwen3 text encoder. Setting that field overrides the Klein default and can trigger downloads of a different text encoder. +> **Important**: For `klein-4b` and `klein-9b`, leave `qwen_text_encoder_model_name_or_path` unset to use the bundled Qwen3 text encoder. Set it only when replacing the bundled Qwen3 text encoder with another Qwen-compatible source. ## Model Overview diff --git a/documentation/quickstart/FLUX2.pt-BR.md b/documentation/quickstart/FLUX2.pt-BR.md index 2d7719f1a..8fac19e0b 100644 --- a/documentation/quickstart/FLUX2.pt-BR.md +++ b/documentation/quickstart/FLUX2.pt-BR.md @@ -25,7 +25,7 @@ Para selecionar uma variante, defina `model_flavour` na sua configuração: } ``` -> **Importante**: Para `klein-4b` e `klein-9b`, deixe `pretrained_text_encoder_model_name_or_path` sem definir, a menos que você realmente queira substituir o encoder Qwen3 incluído. Ao definir esse campo, você sobrescreve o padrão do Klein e pode disparar o download de outro encoder de texto. +> **Importante**: Para `klein-4b` e `klein-9b`, deixe `qwen_text_encoder_model_name_or_path` sem definir para usar o encoder Qwen3 incluído. Defina-o apenas ao substituir o encoder Qwen3 incluído por outra fonte compatível com Qwen. ## Visão geral do modelo diff --git a/documentation/quickstart/FLUX2.zh.md b/documentation/quickstart/FLUX2.zh.md index 027002850..7f17fa61c 100644 --- a/documentation/quickstart/FLUX2.zh.md +++ b/documentation/quickstart/FLUX2.zh.md @@ -25,7 +25,7 @@ FLUX.2 有三个变体: } ``` -> **重要**:对于 `klein-4b` 和 `klein-9b`,除非你明确想替换内置的 Qwen3 文本编码器,否则不要设置 `pretrained_text_encoder_model_name_or_path`。设置这个字段会覆盖 Klein 的默认行为,并可能触发下载其他文本编码器。 +> **重要**:对于 `klein-4b` 和 `klein-9b`,请将 `qwen_text_encoder_model_name_or_path` 留空以使用内置的 Qwen3 文本编码器。仅在需要用另一个兼容 Qwen 的来源替换内置 Qwen3 文本编码器时设置它。 ## 模型概述 diff --git a/documentation/quickstart/index.es.md b/documentation/quickstart/index.es.md index 7605fa6f3..b850bbbed 100644 --- a/documentation/quickstart/index.es.md +++ b/documentation/quickstart/index.es.md @@ -6,65 +6,90 @@ Guías paso a paso para entrenar cada arquitectura de modelo compatible. ### Flow Matching -| Modelo | Parámetros | Guía | -|-------|------------|-------| -| **Flux.1** | 12B | [Guía de Flux.1](FLUX.md) | -| **Flux.2** | 32B | [Guía de Flux.2](FLUX2.md) | -| **Flux Kontext** | 12B | [Guía de Kontext](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Guía de Chroma](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [Guía de SD3](SD3.md) | -| **Auraflow** | 6.8B | [Guía de Auraflow](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Guía de Sana](SANA.md) | -| **Lumina2** | 2B | [Guía de Lumina2](LUMINA2.md) | -| **HiDream** | 17B MoE | [Guía de HiDream](HIDREAM.md) | -| **Z-Image** | - | [Guía de Z-Image](ZIMAGE.md) | -| **Krea2** | - | [Guía de Krea2](KREA2.es.md) | -| **Mage-Flow** | 4B | [Guía de Mage-Flow](MAGEFLOW.es.md) | -| **Boogu-Image** | - | [Guía de Boogu-Image](BOOGU_IMAGE.es.md) | -| **zlab i1** | 3B | [Guía de zlab i1](ZLAB_i1.es.md) | -| **Ideogram 4** | 9B | [Guía de Ideogram 4](IDEOGRAM4.es.md) | -| **ERNIE-Image** | - | [Guía de ERNIE](ERNIE.md) | +| Modelo | Parámetros | Licencia | Permite uso comercial | Guía | +| ------- | ------------ | --- | :---: | ------- | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Aplican condiciones3 | [Guía de Flux.1](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Aplican condiciones4 | [Guía de Flux.2](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | No5 | [Guía de Kontext](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Chroma](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | Aplican condiciones2 | [Guía de SD3](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | Aplican condiciones8 | [Guía de Auraflow](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Sana](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Lumina2](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | Sí | [Guía de HiDream](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Z-Image](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | Sí6 | [Guía de Krea2](KREA2.es.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | Sí | [Guía de Mage-Flow](MAGEFLOW.es.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Boogu-Image](BOOGU_IMAGE.es.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | Sí | [Guía de zlab i1](ZLAB_i1.es.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | No5 | [Guía de Ideogram 4](IDEOGRAM4.es.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de ERNIE](ERNIE.md) | ### DiT / Transformer -| Modelo | Parámetros | Guía | -|-------|------------|-------| -| **PixArt Sigma** | 0.6-0.9B | [Guía de Sigma](SIGMA.md) | -| **Cosmos2** | 2-14B | [Guía de Cosmos2](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Guía de Cosmos3](COSMOS3.es.md) | -| **OmniGen** | 3.8B | [Guía de OmniGen](OMNIGEN.md) | -| **Qwen Image** | 20B | [Guía de Qwen](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [Guía de LongCat](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Guía de Kandinsky](KANDINSKY5_IMAGE.md) | +| Modelo | Parámetros | Licencia | Permite uso comercial | Guía | +| ------- | ------------ | --- | :---: | ------- | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Aplican condiciones1 | [Guía de Sigma](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | Sí9 | [Guía de Cosmos2](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Sí | [Guía de Cosmos3](COSMOS3.es.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | Sí | [Guía de OmniGen](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Qwen](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de LongCat](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | Sí | [Guía de Kandinsky](KANDINSKY5_IMAGE.md) | ### U-Net -| Modelo | Parámetros | Guía | -|-------|------------|-------| -| **Stable Diffusion XL** | 3.5B | [Guía de SDXL](SDXL.md) | -| **Kolors** | 5B | [Guía de Kolors](KOLORS.md) | -| **Stable Cascade** | - | [Guía de Cascade](STABLE_CASCADE_C.md) | +| Modelo | Parámetros | Licencia | Permite uso comercial | Guía | +| ------- | ------------ | --- | :---: | ------- | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Aplican condiciones1 | [Guía de SDXL](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Guía de Kolors](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Guía de Cascade](STABLE_CASCADE_C.md) | ### Edición de imágenes -| Modelo | Guía | -|-------|-------| -| **Qwen Edit** | [Guía de Qwen Edit](QWEN_EDIT.md) | -| **LongCat Edit** | [Guía de LongCat Edit](LONGCAT_EDIT.md) | +| Modelo | Licencia | Permite uso comercial | Guía | +| ------- | --- | :---: | ------- | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Qwen Edit](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de LongCat Edit](LONGCAT_EDIT.md) | ## Modelos de video -| Modelo | Parámetros | Guía | -|-------|------------|-------| -| **Wan Video** | 1.3-14B | [Guía de Wan](WAN.md) | -| **LTX Video** | 5B | [Guía de LTX](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [Guía de LTX Video 2](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Guía de Cosmos3](COSMOS3.es.md) | -| **Hunyuan Video** | 8.3B | [Guía de Hunyuan](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Guía de Sana Video](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Guía de Kandinsky Video](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [Guía de LongCat Video](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [Guía de LongCat Video Edit](LONGCAT_VIDEO_EDIT.md) | +| Modelo | Parámetros | Licencia | Permite uso comercial | Guía | +| ------- | ------------ | --- | :---: | ------- | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Wan](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | Aplican condiciones10 | [Guía de LTX](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | Aplican condiciones10 | [Guía de LTX Video 2](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Sí | [Guía de Cosmos3](COSMOS3.es.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | Aplican condiciones11 | [Guía de Hunyuan](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sí | [Guía de Sana Video](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | Sí | [Guía de Kandinsky Video](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | Sí | [Guía de LongCat Video](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | Sí | [Guía de LongCat Video Edit](LONGCAT_VIDEO_EDIT.md) | + +**Notas de licencia:** El estado de uso comercial cubre pesos del modelo, checkpoints derivados, fine-tunes y uso del modelo alojado. Los derechos sobre salidas generadas pueden diferir; lee el texto de licencia enlazado antes de un despliegue comercial. + +1 Las licencias estilo OpenRAIL suelen permitir uso comercial con restricciones de uso que siguen aplicando al modelo y sus derivados. + +2 La Stability AI Community License está disponible para usuarios que califican por debajo del umbral de ingresos; el uso comercial mayor requiere términos empresariales de Stability. + +3 Flux.1 varía por flavour: Schnell y LibreFlux son Apache-2.0, mientras que Dev, Krea y Kontext usan términos no comerciales de BFL; revisa los metadatos upstream de FluxBooru antes de uso comercial. + +4 Flux.2 varía por flavour: Klein 4B es Apache-2.0, mientras que Dev y Klein 9B usan términos no comerciales de BFL. + +5 Los términos públicos no comerciales no permiten uso comercial de pesos, checkpoints derivados o servicios alojados del modelo sin una licencia separada. + +6 La Krea 2 Community License permite uso comercial bajo su límite de ingresos (menos de $1M anuales) y sus requisitos de seguridad/filtrado; de lo contrario se requiere una licencia empresarial. + +7 Abandonware significa que el proveedor original dejó el modelo atrás y no hay una vía fiable para pedir permiso; cada usuario debe decidir si acepta ese riesgo. + +8 AuraFlow admite flavours upstream Apache-2.0 y un flavour Pony con una licencia personalizada separada; revisa el flavour seleccionado. + +9 La NVIDIA Open Model License permite uso comercial, pero incluye términos de acuerdo, uso aceptable y control de exportación. + +10 LTX Video 0.9.5 usa OpenRAIL-M; LTX Video 2 usa términos comunitarios de LTX con un umbral de ingresos para uso comercial. + +11 La Tencent Hunyuan Community License incluye exclusiones territoriales y un umbral comercial para servicios muy grandes. + ## Modelos de audio diff --git a/documentation/quickstart/index.hi.md b/documentation/quickstart/index.hi.md index 3259ea9b7..56a108939 100644 --- a/documentation/quickstart/index.hi.md +++ b/documentation/quickstart/index.hi.md @@ -6,65 +6,90 @@ ### फ्लो मैचिंग -| मॉडल | पैरामीटर | गाइड | -|-------|------------|-------| -| **Flux.1** | 12B | [Flux.1 गाइड](FLUX.md) | -| **Flux.2** | 32B | [Flux.2 गाइड](FLUX2.md) | -| **Flux Kontext** | 12B | [Kontext गाइड](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Chroma गाइड](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [SD3 गाइड](SD3.md) | -| **Auraflow** | 6.8B | [Auraflow गाइड](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Sana गाइड](SANA.md) | -| **Lumina2** | 2B | [Lumina2 गाइड](LUMINA2.md) | -| **HiDream** | 17B MoE | [HiDream गाइड](HIDREAM.md) | -| **Z-Image** | - | [Z-Image गाइड](ZIMAGE.md) | -| **Krea2** | - | [Krea2 गाइड](KREA2.hi.md) | -| **Mage-Flow** | 4B | [Mage-Flow गाइड](MAGEFLOW.hi.md) | -| **Boogu-Image** | - | [Boogu-Image गाइड](BOOGU_IMAGE.hi.md) | -| **zlab i1** | 3B | [zlab i1 गाइड](ZLAB_i1.hi.md) | -| **Ideogram 4** | 9B | [Ideogram 4 गाइड](IDEOGRAM4.hi.md) | -| **ERNIE-Image** | - | [ERNIE गाइड](ERNIE.md) | +| मॉडल | पैरामीटर | लाइसेंस | व्यावसायिक उपयोग की अनुमति | गाइड | +| ------- | ------------ | --- | :---: | ------- | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | शर्तें लागू3 | [Flux.1 गाइड](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | शर्तें लागू4 | [Flux.2 गाइड](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | नहीं5 | [Kontext गाइड](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Chroma गाइड](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | शर्तें लागू2 | [SD3 गाइड](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | शर्तें लागू8 | [Auraflow गाइड](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Sana गाइड](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Lumina2 गाइड](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | हाँ | [HiDream गाइड](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Z-Image गाइड](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | हाँ6 | [Krea2 गाइड](KREA2.hi.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | हाँ | [Mage-Flow गाइड](MAGEFLOW.hi.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Boogu-Image गाइड](BOOGU_IMAGE.hi.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | हाँ | [zlab i1 गाइड](ZLAB_i1.hi.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | नहीं5 | [Ideogram 4 गाइड](IDEOGRAM4.hi.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [ERNIE गाइड](ERNIE.md) | ### DiT / ट्रांसफ़ॉर्मर -| मॉडल | पैरामीटर | गाइड | -|-------|------------|-------| -| **PixArt Sigma** | 0.6-0.9B | [Sigma गाइड](SIGMA.md) | -| **Cosmos2** | 2-14B | [Cosmos2 गाइड](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Cosmos3 गाइड](COSMOS3.hi.md) | -| **OmniGen** | 3.8B | [OmniGen गाइड](OMNIGEN.md) | -| **Qwen Image** | 20B | [Qwen गाइड](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [LongCat गाइड](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Kandinsky गाइड](KANDINSKY5_IMAGE.md) | +| मॉडल | पैरामीटर | लाइसेंस | व्यावसायिक उपयोग की अनुमति | गाइड | +| ------- | ------------ | --- | :---: | ------- | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | शर्तें लागू1 | [Sigma गाइड](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | हाँ9 | [Cosmos2 गाइड](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | हाँ | [Cosmos3 गाइड](COSMOS3.hi.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | हाँ | [OmniGen गाइड](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Qwen गाइड](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [LongCat गाइड](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | हाँ | [Kandinsky गाइड](KANDINSKY5_IMAGE.md) | ### U‑Net -| मॉडल | पैरामीटर | गाइड | -|-------|------------|-------| -| **Stable Diffusion XL** | 3.5B | [SDXL गाइड](SDXL.md) | -| **Kolors** | 5B | [Kolors गाइड](KOLORS.md) | -| **Stable Cascade** | - | [Cascade गाइड](STABLE_CASCADE_C.md) | +| मॉडल | पैरामीटर | लाइसेंस | व्यावसायिक उपयोग की अनुमति | गाइड | +| ------- | ------------ | --- | :---: | ------- | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | शर्तें लागू1 | [SDXL गाइड](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Kolors गाइड](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Cascade गाइड](STABLE_CASCADE_C.md) | ### छवि संपादन -| मॉडल | गाइड | -|-------|-------| -| **Qwen Edit** | [Qwen Edit गाइड](QWEN_EDIT.md) | -| **LongCat Edit** | [LongCat Edit गाइड](LONGCAT_EDIT.md) | +| मॉडल | लाइसेंस | व्यावसायिक उपयोग की अनुमति | गाइड | +| ------- | --- | :---: | ------- | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Qwen Edit गाइड](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [LongCat Edit गाइड](LONGCAT_EDIT.md) | ## वीडियो मॉडल -| मॉडल | पैरामीटर | गाइड | -|-------|------------|-------| -| **Wan Video** | 1.3-14B | [Wan गाइड](WAN.md) | -| **LTX Video** | 5B | [LTX गाइड](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [LTX Video 2 गाइड](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Cosmos3 गाइड](COSMOS3.hi.md) | -| **Hunyuan Video** | 8.3B | [Hunyuan गाइड](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Sana Video गाइड](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Kandinsky Video गाइड](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [LongCat Video गाइड](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [LongCat Video Edit गाइड](LONGCAT_VIDEO_EDIT.md) | +| मॉडल | पैरामीटर | लाइसेंस | व्यावसायिक उपयोग की अनुमति | गाइड | +| ------- | ------------ | --- | :---: | ------- | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Wan गाइड](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | शर्तें लागू10 | [LTX गाइड](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | शर्तें लागू10 | [LTX Video 2 गाइड](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | हाँ | [Cosmos3 गाइड](COSMOS3.hi.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | शर्तें लागू11 | [Hunyuan गाइड](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | हाँ | [Sana Video गाइड](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | हाँ | [Kandinsky Video गाइड](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | हाँ | [LongCat Video गाइड](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | हाँ | [LongCat Video Edit गाइड](LONGCAT_VIDEO_EDIT.md) | + +**लाइसेंस नोट्स:** व्यावसायिक उपयोग की स्थिति model weights, derivative checkpoints, fine-tunes, और hosted model use को कवर करती है। Generated outputs के अधिकार अलग हो सकते हैं; commercial deployment से पहले linked license text पढ़ें। + +1 OpenRAIL-style licenses आम तौर पर commercial use की अनुमति देती हैं, लेकिन usage restrictions model और derivatives के साथ बनी रहती हैं। + +2 Stability AI Community License revenue threshold से नीचे qualify करने वाले users के लिए उपलब्ध है; बड़े commercial use के लिए Stability enterprise terms चाहिए। + +3 Flux.1 flavour के अनुसार बदलता है: Schnell और LibreFlux Apache-2.0 हैं, जबकि Dev, Krea, और Kontext BFL non-commercial terms का उपयोग करते हैं; commercial use से पहले FluxBooru upstream metadata देखें। + +4 Flux.2 flavour के अनुसार बदलता है: Klein 4B Apache-2.0 है, जबकि Dev और Klein 9B BFL non-commercial terms का उपयोग करते हैं। + +5 Public non-commercial model terms अलग license के बिना weights, derivative checkpoints, या hosted model services के commercial use की अनुमति नहीं देते। + +6 Krea 2 Community License revenue cap (annual revenue $1M से कम) और safety/filtering requirements के तहत commercial use की अनुमति देती है; अन्यथा enterprise license चाहिए। + +7 Abandonware का मतलब है कि original vendor ने model को effectively पीछे छोड़ दिया है और permission लेने का reliable path नहीं है; end users को तय करना होगा कि यह risk स्वीकार्य है या नहीं। + +8 AuraFlow Apache-2.0 upstream flavours और अलग custom license वाले Pony flavour को support करता है; selected flavour देखें। + +9 NVIDIA Open Model License commercial use की अनुमति देती है, लेकिन agreement, acceptable-use, और export-control terms शामिल हैं। + +10 LTX Video 0.9.5 OpenRAIL-M का उपयोग करता है; LTX Video 2 commercial use के लिए revenue threshold वाले LTX community terms का उपयोग करता है। + +11 Tencent Hunyuan Community License में territorial exclusions और बहुत बड़े services के लिए commercial threshold शामिल है। + ## ऑडियो मॉडल diff --git a/documentation/quickstart/index.ja.md b/documentation/quickstart/index.ja.md index a15b0e486..9c94fdfb0 100644 --- a/documentation/quickstart/index.ja.md +++ b/documentation/quickstart/index.ja.md @@ -6,65 +6,90 @@ ### Flow Matching -| モデル | パラメータ | ガイド | -|-------|------------|-------| -| **Flux.1** | 12B | [Flux.1 ガイド](FLUX.md) | -| **Flux.2** | 32B | [Flux.2 ガイド](FLUX2.md) | -| **Flux Kontext** | 12B | [Kontext ガイド](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Chroma ガイド](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [SD3 ガイド](SD3.md) | -| **Auraflow** | 6.8B | [Auraflow ガイド](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Sana ガイド](SANA.md) | -| **Lumina2** | 2B | [Lumina2 ガイド](LUMINA2.md) | -| **HiDream** | 17B MoE | [HiDream ガイド](HIDREAM.md) | -| **Z-Image** | - | [Z-Image ガイド](ZIMAGE.md) | -| **Krea2** | - | [Krea2 ガイド](KREA2.ja.md) | -| **Mage-Flow** | 4B | [Mage-Flow ガイド](MAGEFLOW.ja.md) | -| **Boogu-Image** | - | [Boogu-Image ガイド](BOOGU_IMAGE.ja.md) | -| **zlab i1** | 3B | [zlab i1 ガイド](ZLAB_i1.ja.md) | -| **Ideogram 4** | 9B | [Ideogram 4 ガイド](IDEOGRAM4.ja.md) | -| **ERNIE-Image** | - | [ERNIE ガイド](ERNIE.md) | +| モデル | パラメータ | ライセンス | 商用利用 | ガイド | +| ------- | ------------ | --- | :---: | ------- | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 条件付き3 | [Flux.1 ガイド](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 条件付き4 | [Flux.2 ガイド](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | いいえ5 | [Kontext ガイド](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Chroma ガイド](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | 条件付き2 | [SD3 ガイド](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | 条件付き8 | [Auraflow ガイド](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Sana ガイド](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Lumina2 ガイド](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | はい | [HiDream ガイド](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Z-Image ガイド](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | はい6 | [Krea2 ガイド](KREA2.ja.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | はい | [Mage-Flow ガイド](MAGEFLOW.ja.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Boogu-Image ガイド](BOOGU_IMAGE.ja.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | はい | [zlab i1 ガイド](ZLAB_i1.ja.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | いいえ5 | [Ideogram 4 ガイド](IDEOGRAM4.ja.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [ERNIE ガイド](ERNIE.md) | ### DiT / Transformer -| モデル | パラメータ | ガイド | -|-------|------------|-------| -| **PixArt Sigma** | 0.6-0.9B | [Sigma ガイド](SIGMA.md) | -| **Cosmos2** | 2-14B | [Cosmos2 ガイド](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Cosmos3 ガイド](COSMOS3.ja.md) | -| **OmniGen** | 3.8B | [OmniGen ガイド](OMNIGEN.md) | -| **Qwen Image** | 20B | [Qwen ガイド](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [LongCat ガイド](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Kandinsky ガイド](KANDINSKY5_IMAGE.md) | +| モデル | パラメータ | ライセンス | 商用利用 | ガイド | +| ------- | ------------ | --- | :---: | ------- | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | 条件付き1 | [Sigma ガイド](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | はい9 | [Cosmos2 ガイド](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | はい | [Cosmos3 ガイド](COSMOS3.ja.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | はい | [OmniGen ガイド](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Qwen ガイド](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [LongCat ガイド](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | はい | [Kandinsky ガイド](KANDINSKY5_IMAGE.md) | ### U-Net -| モデル | パラメータ | ガイド | -|-------|------------|-------| -| **Stable Diffusion XL** | 3.5B | [SDXL ガイド](SDXL.md) | -| **Kolors** | 5B | [Kolors ガイド](KOLORS.md) | -| **Stable Cascade** | - | [Cascade ガイド](STABLE_CASCADE_C.md) | +| モデル | パラメータ | ライセンス | 商用利用 | ガイド | +| ------- | ------------ | --- | :---: | ------- | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | 条件付き1 | [SDXL ガイド](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Kolors ガイド](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Cascade ガイド](STABLE_CASCADE_C.md) | ### 画像編集 -| モデル | ガイド | -|-------|-------| -| **Qwen Edit** | [Qwen Edit ガイド](QWEN_EDIT.md) | -| **LongCat Edit** | [LongCat Edit ガイド](LONGCAT_EDIT.md) | +| モデル | ライセンス | 商用利用 | ガイド | +| ------- | --- | :---: | ------- | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Qwen Edit ガイド](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [LongCat Edit ガイド](LONGCAT_EDIT.md) | ## 動画モデル -| モデル | パラメータ | ガイド | -|-------|------------|-------| -| **Wan Video** | 1.3-14B | [Wan ガイド](WAN.md) | -| **LTX Video** | 5B | [LTX ガイド](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [LTX Video 2 ガイド](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Cosmos3 ガイド](COSMOS3.ja.md) | -| **Hunyuan Video** | 8.3B | [Hunyuan ガイド](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Sana Video ガイド](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Kandinsky Video ガイド](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [LongCat Video ガイド](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [LongCat Video Edit ガイド](LONGCAT_VIDEO_EDIT.md) | +| モデル | パラメータ | ライセンス | 商用利用 | ガイド | +| ------- | ------------ | --- | :---: | ------- | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Wan ガイド](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | 条件付き10 | [LTX ガイド](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | 条件付き10 | [LTX Video 2 ガイド](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | はい | [Cosmos3 ガイド](COSMOS3.ja.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | 条件付き11 | [Hunyuan ガイド](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | はい | [Sana Video ガイド](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | はい | [Kandinsky Video ガイド](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | はい | [LongCat Video ガイド](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | はい | [LongCat Video Edit ガイド](LONGCAT_VIDEO_EDIT.md) | + +**ライセンス注記:** 商用利用の表示はモデル重み、派生チェックポイント、fine-tune、ホスト型モデル利用を対象にしています。生成出力の権利は異なる場合があります。商用展開前にリンク先のライセンス本文を確認してください。 + +1 OpenRAIL 系ライセンスは通常、商用利用を許可しますが、モデルと派生物には利用制限が残ります。 + +2 Stability AI Community License は収益しきい値未満の対象ユーザー向けです。より大きな商用利用には Stability のエンタープライズ条件が必要です。 + +3 Flux.1 は flavour により異なります。Schnell と LibreFlux は Apache-2.0、Dev、Krea、Kontext は BFL の非商用条件です。FluxBooru は商用利用前に upstream metadata を確認してください。 + +4 Flux.2 は flavour により異なります。Klein 4B は Apache-2.0、Dev と Klein 9B は BFL の非商用条件です。 + +5 公開されている非商用モデル条件では、別ライセンスなしに重み、派生チェックポイント、ホスト型モデルサービスを商用利用できません。 + +6 Krea 2 Community License は、収益上限(年間 $1M 未満)と安全性/フィルタリング要件を満たす場合に商用利用を許可します。それ以外はエンタープライズライセンスが必要です。 + +7 Abandonware は、元のベンダーが実質的にモデルを放置しており、許可を得る信頼できる経路がないことを意味します。そのリスクを受け入れるかはエンドユーザーの判断です。 + +8 AuraFlow は Apache-2.0 の upstream flavour と、別のカスタムライセンスを持つ Pony flavour をサポートします。選択した flavour を確認してください。 + +9 NVIDIA Open Model License は商用利用を許可しますが、契約、利用許諾ポリシー、輸出管理条件を含みます。 + +10 LTX Video 0.9.5 は OpenRAIL-M、LTX Video 2 は商用利用に収益しきい値がある LTX community terms を使用します。 + +11 Tencent Hunyuan Community License には地域除外と、非常に大規模なサービス向けの商用しきい値があります。 + ## 音声モデル diff --git a/documentation/quickstart/index.md b/documentation/quickstart/index.md index ed755d4c8..df163cdc6 100644 --- a/documentation/quickstart/index.md +++ b/documentation/quickstart/index.md @@ -6,65 +6,90 @@ Step-by-step guides for training each supported model architecture. ### Flow Matching -| Model | Parameters | Guide | -|-------|------------|-------| -| **Flux.1** | 12B | [Flux.1 Guide](FLUX.md) | -| **Flux.2** | 32B | [Flux.2 Guide](FLUX2.md) | -| **Flux Kontext** | 12B | [Kontext Guide](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Chroma Guide](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [SD3 Guide](SD3.md) | -| **Auraflow** | 6.8B | [Auraflow Guide](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Sana Guide](SANA.md) | -| **Lumina2** | 2B | [Lumina2 Guide](LUMINA2.md) | -| **HiDream** | 17B MoE | [HiDream Guide](HIDREAM.md) | -| **Z-Image** | - | [Z-Image Guide](ZIMAGE.md) | -| **Krea2** | - | [Krea2 Guide](KREA2.md) | -| **Mage-Flow** | 4B | [Mage-Flow Guide](MAGEFLOW.md) | -| **Boogu-Image** | - | [Boogu-Image Guide](BOOGU_IMAGE.md) | -| **zlab i1** | 3B | [zlab i1 Guide](ZLAB_i1.md) | -| **Ideogram 4** | 9B | [Ideogram 4 Guide](IDEOGRAM4.md) | -| **ERNIE-Image** | - | [ERNIE Guide](ERNIE.md) | +| Model | Parameters | License | Allows commercial use | Guide | +| ------- | ------------ | --- | :---: | ------- | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Conditions apply3 | [Flux.1 Guide](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Conditions apply4 | [Flux.2 Guide](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | No5 | [Kontext Guide](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Chroma Guide](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | Conditions apply2 | [SD3 Guide](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | Conditions apply8 | [Auraflow Guide](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Sana Guide](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Lumina2 Guide](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | Yes | [HiDream Guide](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Z-Image Guide](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | Yes6 | [Krea2 Guide](KREA2.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | Yes | [Mage-Flow Guide](MAGEFLOW.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Boogu-Image Guide](BOOGU_IMAGE.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | Yes | [zlab i1 Guide](ZLAB_i1.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | No5 | [Ideogram 4 Guide](IDEOGRAM4.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [ERNIE Guide](ERNIE.md) | ### DiT / Transformer -| Model | Parameters | Guide | -|-------|------------|-------| -| **PixArt Sigma** | 0.6-0.9B | [Sigma Guide](SIGMA.md) | -| **Cosmos2** | 2-14B | [Cosmos2 Guide](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Cosmos3 Guide](COSMOS3.md) | -| **OmniGen** | 3.8B | [OmniGen Guide](OMNIGEN.md) | -| **Qwen Image** | 20B | [Qwen Guide](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [LongCat Guide](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Kandinsky Guide](KANDINSKY5_IMAGE.md) | +| Model | Parameters | License | Allows commercial use | Guide | +| ------- | ------------ | --- | :---: | ------- | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Conditions apply1 | [Sigma Guide](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | Yes9 | [Cosmos2 Guide](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Yes | [Cosmos3 Guide](COSMOS3.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | Yes | [OmniGen Guide](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Qwen Guide](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [LongCat Guide](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | Yes | [Kandinsky Guide](KANDINSKY5_IMAGE.md) | ### U-Net -| Model | Parameters | Guide | -|-------|------------|-------| -| **Stable Diffusion XL** | 3.5B | [SDXL Guide](SDXL.md) | -| **Kolors** | 5B | [Kolors Guide](KOLORS.md) | -| **Stable Cascade** | - | [Cascade Guide](STABLE_CASCADE_C.md) | +| Model | Parameters | License | Allows commercial use | Guide | +| ------- | ------------ | --- | :---: | ------- | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Conditions apply1 | [SDXL Guide](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Kolors Guide](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Cascade Guide](STABLE_CASCADE_C.md) | ### Image Editing -| Model | Guide | -|-------|-------| -| **Qwen Edit** | [Qwen Edit Guide](QWEN_EDIT.md) | -| **LongCat Edit** | [LongCat Edit Guide](LONGCAT_EDIT.md) | +| Model | License | Allows commercial use | Guide | +| ------- | --- | :---: | ------- | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Qwen Edit Guide](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [LongCat Edit Guide](LONGCAT_EDIT.md) | ## Video Models -| Model | Parameters | Guide | -|-------|------------|-------| -| **Wan Video** | 1.3-14B | [Wan Guide](WAN.md) | -| **LTX Video** | 5B | [LTX Guide](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [LTX Video 2 Guide](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Cosmos3 Guide](COSMOS3.md) | -| **Hunyuan Video** | 8.3B | [Hunyuan Guide](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Sana Video Guide](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Kandinsky Video Guide](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [LongCat Video Guide](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [LongCat Video Edit Guide](LONGCAT_VIDEO_EDIT.md) | +| Model | Parameters | License | Allows commercial use | Guide | +| ------- | ------------ | --- | :---: | ------- | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Wan Guide](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | Conditions apply10 | [LTX Guide](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | Conditions apply10 | [LTX Video 2 Guide](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Yes | [Cosmos3 Guide](COSMOS3.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | Conditions apply11 | [Hunyuan Guide](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Yes | [Sana Video Guide](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | Yes | [Kandinsky Video Guide](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | Yes | [LongCat Video Guide](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | Yes | [LongCat Video Edit Guide](LONGCAT_VIDEO_EDIT.md) | + +**License notes:** The commercial-use status covers model weights, derivative checkpoints, fine-tunes, and hosted model use. Generated-output rights can differ; read the linked license text before commercial deployment. + +1 OpenRAIL-style licenses generally permit commercial use with usage restrictions that remain attached to the model and derivatives. + +2 Stability AI Community License is available for qualifying users below the revenue threshold; larger commercial use needs Stability enterprise terms. + +3 Flux.1 varies by flavour: Schnell and LibreFlux are Apache-2.0, while Dev, Krea, and Kontext use BFL non-commercial terms; review FluxBooru upstream metadata before commercial use. + +4 Flux.2 varies by flavour: Klein 4B is Apache-2.0, while Dev and Klein 9B use BFL non-commercial terms. + +5 Public non-commercial model terms do not permit commercial use of weights, derivative checkpoints, or hosted model services without a separate license. + +6 Krea 2 Community License permits commercial use under its revenue cap (under $1M annual revenue) and safety/filtering requirements; otherwise an enterprise license is required. + +7 Abandonware means the original vendor has effectively left the model behind and there is no reliable permission path; end users must decide whether that risk is acceptable. + +8 AuraFlow supports Apache-2.0 upstream flavours and a Pony flavour with a separate custom license; check the selected flavour. + +9 NVIDIA Open Model License permits commercial use but includes agreement, acceptable-use, and export-control terms. + +10 LTX Video 0.9.5 uses OpenRAIL-M; LTX Video 2 uses LTX community terms with a revenue threshold for commercial use. + +11 Tencent Hunyuan Community License includes territorial exclusions and a commercial threshold for very large services. + ## Audio Models diff --git a/documentation/quickstart/index.pt-BR.md b/documentation/quickstart/index.pt-BR.md index 09d81d9ef..1f37312d8 100644 --- a/documentation/quickstart/index.pt-BR.md +++ b/documentation/quickstart/index.pt-BR.md @@ -6,65 +6,90 @@ Guias passo a passo para treinar cada arquitetura de modelo suportada. ### Flow Matching -| Modelo | Parâmetros | Guia | -|-------|------------|------| -| **Flux.1** | 12B | [Guia Flux.1](FLUX.md) | -| **Flux.2** | 32B | [Guia Flux.2](FLUX2.md) | -| **Flux Kontext** | 12B | [Guia Kontext](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Guia Chroma](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [Guia SD3](SD3.md) | -| **Auraflow** | 6.8B | [Guia Auraflow](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Guia Sana](SANA.md) | -| **Lumina2** | 2B | [Guia Lumina2](LUMINA2.md) | -| **HiDream** | 17B MoE | [Guia HiDream](HIDREAM.md) | -| **Z-Image** | - | [Guia Z-Image](ZIMAGE.md) | -| **Krea2** | - | [Guia Krea2](KREA2.pt-BR.md) | -| **Mage-Flow** | 4B | [Guia Mage-Flow](MAGEFLOW.pt-BR.md) | -| **Boogu-Image** | - | [Guia Boogu-Image](BOOGU_IMAGE.pt-BR.md) | -| **zlab i1** | 3B | [Guia zlab i1](ZLAB_i1.pt-BR.md) | -| **Ideogram 4** | 9B | [Guia Ideogram 4](IDEOGRAM4.pt-BR.md) | -| **ERNIE-Image** | - | [Guia ERNIE](ERNIE.md) | +| Modelo | Parâmetros | Licença | Permite uso comercial | Guia | +| ------- | ------------ | --- | :---: | ------ | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Condições aplicáveis3 | [Guia Flux.1](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Condições aplicáveis4 | [Guia Flux.2](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | Não5 | [Guia Kontext](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Chroma](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | Condições aplicáveis2 | [Guia SD3](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | Condições aplicáveis8 | [Guia Auraflow](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Sana](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Lumina2](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | Sim | [Guia HiDream](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Z-Image](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | Sim6 | [Guia Krea2](KREA2.pt-BR.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | Sim | [Guia Mage-Flow](MAGEFLOW.pt-BR.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Boogu-Image](BOOGU_IMAGE.pt-BR.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | Sim | [Guia zlab i1](ZLAB_i1.pt-BR.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | Não5 | [Guia Ideogram 4](IDEOGRAM4.pt-BR.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia ERNIE](ERNIE.md) | ### DiT / Transformador -| Modelo | Parâmetros | Guia | -|-------|------------|------| -| **PixArt Sigma** | 0.6-0.9B | [Guia Sigma](SIGMA.md) | -| **Cosmos2** | 2-14B | [Guia Cosmos2](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Guia Cosmos3](COSMOS3.pt-BR.md) | -| **OmniGen** | 3.8B | [Guia OmniGen](OMNIGEN.md) | -| **Qwen Image** | 20B | [Guia Qwen](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [Guia LongCat](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Guia Kandinsky](KANDINSKY5_IMAGE.md) | +| Modelo | Parâmetros | Licença | Permite uso comercial | Guia | +| ------- | ------------ | --- | :---: | ------ | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Condições aplicáveis1 | [Guia Sigma](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | Sim9 | [Guia Cosmos2](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Sim | [Guia Cosmos3](COSMOS3.pt-BR.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | Sim | [Guia OmniGen](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Qwen](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia LongCat](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | Sim | [Guia Kandinsky](KANDINSKY5_IMAGE.md) | ### U-Net -| Modelo | Parâmetros | Guia | -|-------|------------|------| -| **Stable Diffusion XL** | 3.5B | [Guia SDXL](SDXL.md) | -| **Kolors** | 5B | [Guia Kolors](KOLORS.md) | -| **Stable Cascade** | - | [Guia Cascade](STABLE_CASCADE_C.md) | +| Modelo | Parâmetros | Licença | Permite uso comercial | Guia | +| ------- | ------------ | --- | :---: | ------ | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | Condições aplicáveis1 | [Guia SDXL](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Guia Kolors](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Guia Cascade](STABLE_CASCADE_C.md) | ### Edição de Imagem -| Modelo | Guia | -|-------|------| -| **Qwen Edit** | [Guia Qwen Edit](QWEN_EDIT.md) | -| **LongCat Edit** | [Guia LongCat Edit](LONGCAT_EDIT.md) | +| Modelo | Licença | Permite uso comercial | Guia | +| ------- | --- | :---: | ------ | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Qwen Edit](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia LongCat Edit](LONGCAT_EDIT.md) | ## Modelos de Vídeo -| Modelo | Parâmetros | Guia | -|-------|------------|------| -| **Wan Video** | 1.3-14B | [Guia Wan](WAN.md) | -| **LTX Video** | 5B | [Guia LTX](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [Guia LTX Video 2](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Guia Cosmos3](COSMOS3.pt-BR.md) | -| **Hunyuan Video** | 8.3B | [Guia Hunyuan](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Guia Sana Video](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Guia Kandinsky Video](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [Guia LongCat Video](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [Guia LongCat Video Edit](LONGCAT_VIDEO_EDIT.md) | +| Modelo | Parâmetros | Licença | Permite uso comercial | Guia | +| ------- | ------------ | --- | :---: | ------ | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Wan](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | Condições aplicáveis10 | [Guia LTX](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | Condições aplicáveis10 | [Guia LTX Video 2](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | Sim | [Guia Cosmos3](COSMOS3.pt-BR.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | Condições aplicáveis11 | [Guia Hunyuan](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | Sim | [Guia Sana Video](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | Sim | [Guia Kandinsky Video](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | Sim | [Guia LongCat Video](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | Sim | [Guia LongCat Video Edit](LONGCAT_VIDEO_EDIT.md) | + +**Notas de licença:** O status de uso comercial cobre pesos do modelo, checkpoints derivados, fine-tunes e uso de modelo hospedado. Os direitos sobre saídas geradas podem diferir; leia o texto da licença vinculada antes de uma implantação comercial. + +1 Licenças estilo OpenRAIL geralmente permitem uso comercial com restrições de uso que continuam aplicáveis ao modelo e seus derivados. + +2 A Stability AI Community License está disponível para usuários qualificados abaixo do limite de receita; uso comercial maior exige termos empresariais da Stability. + +3 Flux.1 varia por flavour: Schnell e LibreFlux são Apache-2.0, enquanto Dev, Krea e Kontext usam termos não comerciais da BFL; revise os metadados upstream do FluxBooru antes de uso comercial. + +4 Flux.2 varia por flavour: Klein 4B é Apache-2.0, enquanto Dev e Klein 9B usam termos não comerciais da BFL. + +5 Termos públicos de modelo não comercial não permitem uso comercial de pesos, checkpoints derivados ou serviços hospedados do modelo sem uma licença separada. + +6 A Krea 2 Community License permite uso comercial sob seu limite de receita (menos de $1M anuais) e requisitos de segurança/filtragem; caso contrário, é necessária uma licença empresarial. + +7 Abandonware significa que o fornecedor original efetivamente deixou o modelo para trás e não há um caminho confiável para pedir permissão; cada usuário deve decidir se aceita esse risco. + +8 AuraFlow aceita flavours upstream Apache-2.0 e um flavour Pony com uma licença personalizada separada; confira o flavour selecionado. + +9 A NVIDIA Open Model License permite uso comercial, mas inclui termos de contrato, uso aceitável e controle de exportação. + +10 LTX Video 0.9.5 usa OpenRAIL-M; LTX Video 2 usa termos comunitários da LTX com limite de receita para uso comercial. + +11 A Tencent Hunyuan Community License inclui exclusões territoriais e um limite comercial para serviços muito grandes. + ## Modelos de Áudio diff --git a/documentation/quickstart/index.zh.md b/documentation/quickstart/index.zh.md index 594fb5859..dc5271eb9 100644 --- a/documentation/quickstart/index.zh.md +++ b/documentation/quickstart/index.zh.md @@ -6,65 +6,90 @@ ### 流匹配 -| 模型 | 参数 | 指南 | -|-------|------------|-------| -| **Flux.1** | 12B | [Flux.1 指南](FLUX.md) | -| **Flux.2** | 32B | [Flux.2 指南](FLUX2.md) | -| **Flux Kontext** | 12B | [Kontext 指南](FLUX_KONTEXT.md) | -| **Chroma** | 8.9B | [Chroma 指南](CHROMA.md) | -| **Stable Diffusion 3** | 2-8B | [SD3 指南](SD3.md) | -| **Auraflow** | 6.8B | [Auraflow 指南](AURAFLOW.md) | -| **Sana** | 0.6-4.8B | [Sana 指南](SANA.md) | -| **Lumina2** | 2B | [Lumina2 指南](LUMINA2.md) | -| **HiDream** | 17B MoE | [HiDream 指南](HIDREAM.md) | -| **Z-Image** | - | [Z-Image 指南](ZIMAGE.md) | -| **Krea2** | - | [Krea2 指南](KREA2.zh.md) | -| **Mage-Flow** | 4B | [Mage-Flow 指南](MAGEFLOW.zh.md) | -| **Boogu-Image** | - | [Boogu-Image 指南](BOOGU_IMAGE.zh.md) | -| **zlab i1** | 3B | [zlab i1 指南](ZLAB_i1.zh.md) | -| **Ideogram 4** | 9B | [Ideogram 4 指南](IDEOGRAM4.zh.md) | -| **ERNIE-Image** | - | [ERNIE 指南](ERNIE.md) | +| 模型 | 参数 | 许可证 | 允许商用 | 指南 | +| ------- | ------------ | --- | :---: | ------- | +| **Flux.1** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 有条件3 | [Flux.1 指南](FLUX.md) | +| **Flux.2** | 32B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) / [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 有条件4 | [Flux.2 指南](FLUX2.md) | +| **Flux Kontext** | 12B | [BFL Non-Commercial](https://bfl.ai/legal/non-commercial-license-terms) | 否5 | [Kontext 指南](FLUX_KONTEXT.md) | +| **Chroma** | 8.9B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Chroma 指南](CHROMA.md) | +| **Stable Diffusion 3** | 2-8B | [Stability AI Community](https://stability.ai/license) | 有条件2 | [SD3 指南](SD3.md) | +| **Auraflow** | 6.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) / [Pony License](https://huggingface.co/purplesmartai/pony-v7-base/blob/main/LICENSE) | 有条件8 | [Auraflow 指南](AURAFLOW.md) | +| **Sana** | 0.6-4.8B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Sana 指南](SANA.md) | +| **Lumina2** | 2B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Lumina2 指南](LUMINA2.md) | +| **HiDream** | 17B MoE | [MIT](https://opensource.org/license/mit) | 是 | [HiDream 指南](HIDREAM.md) | +| **Z-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Z-Image 指南](ZIMAGE.md) | +| **Krea2** | - | [Krea 2 Community](https://www.krea.ai/krea-2-licensing) | 是6 | [Krea2 指南](KREA2.zh.md) | +| **Mage-Flow** | 4B | [MIT](https://opensource.org/license/mit) | 是 | [Mage-Flow 指南](MAGEFLOW.zh.md) | +| **Boogu-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Boogu-Image 指南](BOOGU_IMAGE.zh.md) | +| **zlab i1** | 3B | [MIT](https://opensource.org/license/mit) | 是 | [zlab i1 指南](ZLAB_i1.zh.md) | +| **Ideogram 4** | 9B | [Ideogram 4 Non-Commercial](https://huggingface.co/ideogram-ai/ideogram-4-nf4/blob/main/LICENSE.md) | 否5 | [Ideogram 4 指南](IDEOGRAM4.zh.md) | +| **ERNIE-Image** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [ERNIE 指南](ERNIE.md) | ### DiT / Transformer -| 模型 | 参数 | 指南 | -|-------|------------|-------| -| **PixArt Sigma** | 0.6-0.9B | [Sigma 指南](SIGMA.md) | -| **Cosmos2** | 2-14B | [Cosmos2 指南](COSMOS2IMAGE.md) | -| **Cosmos3** | 4-65B | [Cosmos3 指南](COSMOS3.zh.md) | -| **OmniGen** | 3.8B | [OmniGen 指南](OMNIGEN.md) | -| **Qwen Image** | 20B | [Qwen 指南](QWEN_IMAGE.md) | -| **LongCat Image** | 6B | [LongCat 指南](LONGCAT_IMAGE.md) | -| **Kandinsky 5** | - | [Kandinsky 指南](KANDINSKY5_IMAGE.md) | +| 模型 | 参数 | 许可证 | 允许商用 | 指南 | +| ------- | ------------ | --- | :---: | ------- | +| **PixArt Sigma** | 0.6-0.9B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | 有条件1 | [Sigma 指南](SIGMA.md) | +| **Cosmos2** | 2-14B | [NVIDIA Open Model License](https://www.nvidia.com/en-us/agreements/enterprise-software/nvidia-open-model-license/) | 是9 | [Cosmos2 指南](COSMOS2IMAGE.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | 是 | [Cosmos3 指南](COSMOS3.zh.md) | +| **OmniGen** | 3.8B | [MIT](https://opensource.org/license/mit) | 是 | [OmniGen 指南](OMNIGEN.md) | +| **Qwen Image** | 20B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Qwen 指南](QWEN_IMAGE.md) | +| **LongCat Image** | 6B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [LongCat 指南](LONGCAT_IMAGE.md) | +| **Kandinsky 5** | - | [MIT](https://opensource.org/license/mit) | 是 | [Kandinsky 指南](KANDINSKY5_IMAGE.md) | ### U-Net -| 模型 | 参数 | 指南 | -|-------|------------|-------| -| **Stable Diffusion XL** | 3.5B | [SDXL 指南](SDXL.md) | -| **Kolors** | 5B | [Kolors 指南](KOLORS.md) | -| **Stable Cascade** | - | [Cascade 指南](STABLE_CASCADE_C.md) | +| 模型 | 参数 | 许可证 | 允许商用 | 指南 | +| ------- | ------------ | --- | :---: | ------- | +| **Stable Diffusion XL** | 3.5B | [OpenRAIL++](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) | 有条件1 | [SDXL 指南](SDXL.md) | +| **Kolors** | 5B | [Kwai Kolors License](https://huggingface.co/terminusresearch/kwai-kolors-1.0/blob/main/MODEL_LICENSE) | Abandonware7 | [Kolors 指南](KOLORS.md) | +| **Stable Cascade** | - | [Stable Cascade NC Community](https://huggingface.co/stabilityai/stable-cascade/blob/main/LICENSE) | Abandonware7 | [Cascade 指南](STABLE_CASCADE_C.md) | ### 图像编辑 -| 模型 | 指南 | -|-------|-------| -| **Qwen Edit** | [Qwen Edit 指南](QWEN_EDIT.md) | -| **LongCat Edit** | [LongCat Edit 指南](LONGCAT_EDIT.md) | +| 模型 | 许可证 | 允许商用 | 指南 | +| ------- | --- | :---: | ------- | +| **Qwen Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Qwen Edit 指南](QWEN_EDIT.md) | +| **LongCat Edit** | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [LongCat Edit 指南](LONGCAT_EDIT.md) | ## 视频模型 -| 模型 | 参数 | 指南 | -|-------|------------|-------| -| **Wan Video** | 1.3-14B | [Wan 指南](WAN.md) | -| **LTX Video** | 5B | [LTX 指南](LTXVIDEO.md) | -| **LTX Video 2** | 19B | [LTX Video 2 指南](LTXVIDEO2.md) | -| **Cosmos3** | 4-65B | [Cosmos3 指南](COSMOS3.zh.md) | -| **Hunyuan Video** | 8.3B | [Hunyuan 指南](HUNYUANVIDEO.md) | -| **Sana Video** | - | [Sana Video 指南](SANAVIDEO.md) | -| **Kandinsky 5 Video** | - | [Kandinsky Video 指南](KANDINSKY5_VIDEO.md) | -| **LongCat Video** | - | [LongCat Video 指南](LONGCAT_VIDEO.md) | -| **LongCat Video Edit** | - | [LongCat Video Edit 指南](LONGCAT_VIDEO_EDIT.md) | +| 模型 | 参数 | 许可证 | 允许商用 | 指南 | +| ------- | ------------ | --- | :---: | ------- | +| **Wan Video** | 1.3-14B | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Wan 指南](WAN.md) | +| **LTX Video** | 5B | [LTX Video OpenRAIL-M](https://huggingface.co/Lightricks/LTX-Video-0.9.5/blob/main/ltx-video-2b-v0.9.5.license.txt) | 有条件10 | [LTX 指南](LTXVIDEO.md) | +| **LTX Video 2** | 19B | [LTX-2 Community](https://ltx.io/model/license) | 有条件10 | [LTX Video 2 指南](LTXVIDEO2.md) | +| **Cosmos3** | 4-65B | [OpenMDW 1.1](https://github.com/OpenMDW/openmdw/blob/main/1.1/LICENSE.OpenMDW-1.1) | 是 | [Cosmos3 指南](COSMOS3.zh.md) | +| **Hunyuan Video** | 8.3B | [Tencent Hunyuan Community](https://huggingface.co/tencent/HunyuanVideo-1.5/blob/main/LICENSE) | 有条件11 | [Hunyuan 指南](HUNYUANVIDEO.md) | +| **Sana Video** | - | [Apache-2.0](https://www.apache.org/licenses/LICENSE-2.0) | 是 | [Sana Video 指南](SANAVIDEO.md) | +| **Kandinsky 5 Video** | - | [MIT](https://opensource.org/license/mit) | 是 | [Kandinsky Video 指南](KANDINSKY5_VIDEO.md) | +| **LongCat Video** | - | [MIT](https://opensource.org/license/mit) | 是 | [LongCat Video 指南](LONGCAT_VIDEO.md) | +| **LongCat Video Edit** | - | [MIT](https://opensource.org/license/mit) | 是 | [LongCat Video Edit 指南](LONGCAT_VIDEO_EDIT.md) | + +**许可证说明:** 商用状态覆盖模型权重、派生 checkpoint、fine-tune 和托管模型使用。生成输出的权利可能不同;商业部署前请以链接的许可证正文为准。 + +1 OpenRAIL 风格许可证通常允许商用,但使用限制仍适用于模型和派生物。 + +2 Stability AI Community License 适用于低于收入门槛且符合条件的用户;更大规模商用需要 Stability 企业条款。 + +3 Flux.1 随 flavour 不同而不同:Schnell 和 LibreFlux 为 Apache-2.0,Dev、Krea 和 Kontext 使用 BFL 非商用条款;FluxBooru 商用前请检查 upstream metadata。 + +4 Flux.2 随 flavour 不同而不同:Klein 4B 为 Apache-2.0,Dev 和 Klein 9B 使用 BFL 非商用条款。 + +5 公开的非商用模型条款不允许在没有单独许可证的情况下商用权重、派生 checkpoint 或托管模型服务。 + +6 Krea 2 Community License 在满足收入上限(年收入低于 $1M)和安全/过滤要求时允许商用;否则需要企业许可证。 + +7 Abandonware 表示原供应商实际上已放弃维护该模型,且没有可靠的许可申请路径;最终用户需要自行判断是否接受该风险。 + +8 AuraFlow 支持 Apache-2.0 upstream flavour,以及带有单独自定义许可证的 Pony flavour;请检查所选 flavour。 + +9 NVIDIA Open Model License 允许商用,但包含协议、可接受使用和出口管制条款。 + +10 LTX Video 0.9.5 使用 OpenRAIL-M;LTX Video 2 使用带商用收入门槛的 LTX community terms。 + +11 Tencent Hunyuan Community License 包含地域排除,以及针对超大规模服务的商用门槛。 + ## 音频模型 diff --git a/mkdocs.yml b/mkdocs.yml index 787f36654..8a4980ebb 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -357,6 +357,7 @@ nav: - Metal Flash Attention: experimental/METAL_FLASH_ATTENTION.md - Prompt2Effect: experimental/PROMPT2EFFECT.md - Self-Flow: experimental/SELF_FLOW.md + - Segmented Checkpointing: experimental/SEGMENTED_CHECKPOINTING.md - Scheduled Sampling: experimental/SCHEDULED_SAMPLING.md - T-LoRA: experimental/T_LORA.md - Unsloth-Style Checkpointing: experimental/UNSLOTH_CHECKPOINTING.md diff --git a/setup.py b/setup.py index 7a9f0f5bd..344482a1b 100644 --- a/setup.py +++ b/setup.py @@ -317,7 +317,7 @@ def _collect_package_files(*directories: str): "hangul-romanize>=0.1.0", "optimum-quanto>=0.2.7", "lycoris-lora>=3.4.0", - "kernels>=0.12.3,<0.13", + "kernels>=0.15.2,<0.16.0", "torch-optimi>=0.2.1", "librosa>=0.10.2", "loguru>=0.7.2", diff --git a/simpletuner/__init__.py b/simpletuner/__init__.py index 54fd31a58..09edbf11d 100644 --- a/simpletuner/__init__.py +++ b/simpletuner/__init__.py @@ -126,4 +126,4 @@ def _suppress_swigvarlink(message, *args, **kwargs): warnings.warn = _suppress_swigvarlink -__version__ = "4.5.2" +__version__ = "4.6.0" diff --git a/simpletuner/examples/acestep-audio.json b/simpletuner/examples/acestep-audio.json index 365ba2e4d..312b10dd1 100644 --- a/simpletuner/examples/acestep-audio.json +++ b/simpletuner/examples/acestep-audio.json @@ -6,6 +6,11 @@ "dataset_name": "Yi3852/ACEStep-Songs", "metadata_backend": "huggingface", "caption_strategy": "huggingface", + "audio": { + "bucket_strategy": "duration", + "duration_interval": 3.0, + "max_duration_seconds": 30 + }, "cache_dir_vae": "cache/vae/{model_family}/acestep-demo-data" }, { diff --git a/simpletuner/examples/anima.peft-lora/config.json b/simpletuner/examples/anima.peft-lora/config.json new file mode 100644 index 000000000..32ab62704 --- /dev/null +++ b/simpletuner/examples/anima.peft-lora/config.json @@ -0,0 +1,56 @@ +{ + "base_model_precision": "no_change", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 1000, + "checkpoints_total_limit": 6, + "compress_disk_cache": false, + "data_backend_config": "config/examples/anima.peft-lora/multidatabackend.json", + "disable_benchmark": true, + "flow_schedule_shift": 3.0, + "gradient_accumulation_steps": 1, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-anima-peft-lora", + "strict_epoch_limit": false, + "learning_rate": 1e-5, + "lora_alpha": 64, + "lora_rank": 64, + "lora_type": "standard", + "loss_type": "l2", + "lr_scheduler": "constant_with_warmup", + "lr_warmup_steps": 250, + "max_grad_norm": 0.01, + "max_train_steps": 1000, + "mixed_precision": "bf16", + "model_family": "anima", + "model_flavour": "base-v1.0", + "model_type": "lora", + "num_eval_images": 1, + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/anima.peft-lora", + "pretrained_model_name_or_path": "circlestone-labs/Anima-Base-v1.0-Diffusers", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "report_to": "tensorboard", + "resolution": 1024, + "resolution_type": "pixel_area", + "resume_from_checkpoint": "latest", + "seed": 42, + "skip_file_discovery": false, + "tracker_project_name": "lora-training", + "tracker_run_name": "anima-peft-lora", + "train_batch_size": 1, + "use_ema": false, + "user_prompt_library": "config/examples/anima.peft-lora/user_prompt_library.json", + "vae_batch_size": 1, + "validation_disable_unconditional": true, + "validation_guidance": 4.0, + "validation_guidance_rescale": 0.0, + "validation_negative_prompt": "worst quality, low quality, score_1, score_2, score_3, blurry, cropped, artist name, signature", + "validation_num_inference_steps": 20, + "validation_prompt": "masterpiece, best quality, score_7, safe, anime portrait of a young woman with blue hair wearing a white jacket, clean line art, detailed eyes, soft daylight", + "validation_prompt_library": false, + "validation_resolution": "1024x1024", + "validation_seed": 42, + "validation_steps": 1000 +} diff --git a/simpletuner/examples/anima.peft-lora/multidatabackend.json b/simpletuner/examples/anima.peft-lora/multidatabackend.json new file mode 100644 index 000000000..354f0c0d1 --- /dev/null +++ b/simpletuner/examples/anima.peft-lora/multidatabackend.json @@ -0,0 +1,35 @@ +[ + { + "id": "anime-caption-hf", + "type": "huggingface", + "dataset_name": "Dhiraj45/Anime-Caption", + "caption_strategy": "huggingface", + "metadata_backend": "huggingface", + "image_column": "Image", + "caption_column": "Caption", + "huggingface": { + "dataset_name": "Dhiraj45/Anime-Caption", + "image_column": "Image", + "caption_column": "Caption", + "file_extension": "jpg" + }, + "crop": true, + "crop_style": "center", + "crop_aspect": "square", + "minimum_image_size": 512, + "maximum_image_size": 1536, + "target_downsample_size": 1024, + "resolution": 1024, + "resolution_type": "pixel_area", + "cache_dir_vae": "cache/vae/anima-example", + "text_embeds": "anima-example-text", + "repeats": 1 + }, + { + "id": "anima-example-text", + "dataset_type": "text_embeds", + "default": true, + "type": "local", + "cache_dir": "cache/text/anima-example" + } +] diff --git a/simpletuner/examples/anima.peft-lora/user_prompt_library.json b/simpletuner/examples/anima.peft-lora/user_prompt_library.json new file mode 100644 index 000000000..26db125f9 --- /dev/null +++ b/simpletuner/examples/anima.peft-lora/user_prompt_library.json @@ -0,0 +1,6 @@ +{ + "portrait": "masterpiece, best quality, score_7, safe, anime portrait of a young woman with blue hair wearing a white jacket, clean line art, detailed eyes, soft daylight", + "character_full_body": "masterpiece, best quality, score_7, safe, full body anime character design, red scarf, streetwear, simple studio background, sharp silhouette", + "environment": "masterpiece, best quality, score_7, safe, cozy fantasy cafe interior, warm lantern light, detailed anime background art, inviting atmosphere", + "object_scene": "masterpiece, best quality, score_7, safe, small robot mascot tending flowers in a garden, watercolor anime style, clear subject, bright colors" +} diff --git a/simpletuner/examples/cascade-stage-c.lycoris-lokr/config.json b/simpletuner/examples/cascade-stage-c.lycoris-lokr/config.json index 9a77c78c6..b2b02665c 100644 --- a/simpletuner/examples/cascade-stage-c.lycoris-lokr/config.json +++ b/simpletuner/examples/cascade-stage-c.lycoris-lokr/config.json @@ -14,11 +14,12 @@ "lr_scheduler": "constant", "lycoris_config": "config/examples/cascade-stage-c.lycoris-lokr/lycoris_config.json", "max_train_steps": 100, + "mixed_precision": "no", "model_family": "stable_cascade", "model_type": "lora", "num_eval_images": 25, "num_train_epochs": 0, - "optimizer": "adamw_bf16", + "optimizer": "torch-adamw", "output_dir": "output/examples/cascade-stage-c.lycoris-lokr", "push_checkpoints_to_hub": false, "push_to_hub": false, diff --git a/simpletuner/examples/cosmos2image.lycoris-lokr/config.json b/simpletuner/examples/cosmos2image.lycoris-lokr/config.json index ef2d74052..7acb14592 100644 --- a/simpletuner/examples/cosmos2image.lycoris-lokr/config.json +++ b/simpletuner/examples/cosmos2image.lycoris-lokr/config.json @@ -15,7 +15,7 @@ "lr_scheduler": "constant", "lycoris_config": "config/examples/cosmos2image.lycoris-lokr/lycoris_config.json", "max_train_steps": 100, - "model_family": "cosmos2image", + "model_family": "cosmos", "model_type": "lora", "num_train_epochs": 0, "optimizer": "adamw_bf16", diff --git a/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/config.json b/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/config.json new file mode 100644 index 000000000..51dc61fdb --- /dev/null +++ b/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/config.json @@ -0,0 +1,38 @@ +{ + "base_model_precision": "no_change", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 50, + "checkpoints_total_limit": 5, + "data_backend_config": "config/examples/deepfloyd-if-i-medium.peft-lora/multidatabackend.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-deepfloyd-if-i-medium-peft-lora", + "learning_rate": 1e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant", + "max_train_steps": 100, + "mixed_precision": "bf16", + "model_family": "deepfloyd", + "model_flavour": "i-medium-400m", + "pretrained_model_name_or_path": "DeepFloyd/IF-I-M-v1.0", + "model_type": "lora", + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/deepfloyd-if-i-medium.peft-lora", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "report_to": "none", + "resolution": 512, + "resolution_type": "pixel", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "validation_guidance": 7.0, + "validation_num_inference_steps": 20, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "512x512", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/multidatabackend.json b/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/multidatabackend.json new file mode 100644 index 000000000..87440bd6b --- /dev/null +++ b/simpletuner/examples/deepfloyd-if-i-medium.peft-lora/multidatabackend.json @@ -0,0 +1,25 @@ +[ + { + "id": "deepfloyd-local-512", + "type": "local", + "instance_data_dir": "datasets/ideogram-subject", + "crop": true, + "crop_style": "random", + "crop_aspect": "square", + "minimum_image_size": 128, + "maximum_image_size": 512, + "target_downsample_size": 512, + "resolution": 512, + "resolution_type": "pixel", + "metadata_backend": "discovery", + "caption_strategy": "filename", + "repeats": 2 + }, + { + "id": "deepfloyd-text-cache", + "dataset_type": "text_embeds", + "default": true, + "type": "local", + "cache_dir": "cache/text/deepfloyd" + } +] diff --git a/simpletuner/examples/heartmula.peft-lora/config.json b/simpletuner/examples/heartmula.peft-lora/config.json index b0f97584f..c8c9ee810 100644 --- a/simpletuner/examples/heartmula.peft-lora/config.json +++ b/simpletuner/examples/heartmula.peft-lora/config.json @@ -17,7 +17,6 @@ "optimizer": "adamw_bf16", "output_dir": "output/examples/heartmula.peft-lora", "report_to": "none", - "resolution": 512, "seed": 42, "tracker_project_name": "lora-training", "tracker_run_name": "example-training-run", @@ -29,7 +28,7 @@ "validation_disable_unconditional": true, "validation_guidance": 3, "validation_lyrics": "I'm a little tea-pot, short and stout. Here is my handle, here is my spout. warming up for a little fun, when I get all steamed up, hear me shout. don't you stress out, just let it all out.", - "validation_prompt": "🟫 is holding a sign that says hello world from heartmula", + "validation_prompt": "bright synth pop, steady percussion, clean vocal", "validation_steps": 50, "validation_num_inference_steps": 80, "validation_seed": 42 diff --git a/simpletuner/examples/hidream.peft-lora/config.json b/simpletuner/examples/hidream.peft-lora/config.json index 3ccfdfb2b..71e289c7c 100644 --- a/simpletuner/examples/hidream.peft-lora/config.json +++ b/simpletuner/examples/hidream.peft-lora/config.json @@ -65,5 +65,5 @@ "validation_prompt_library": false, "validation_resolution": "512x512", "validation_seed": 42, - "validation_steps": 10, + "validation_steps": 10 } diff --git a/simpletuner/examples/kandinsky5-image-6b-i2i.lycoris-lokr/config.json b/simpletuner/examples/kandinsky5-image-6b-i2i.lycoris-lokr/config.json index a807a22ef..b7f43a7d2 100644 --- a/simpletuner/examples/kandinsky5-image-6b-i2i.lycoris-lokr/config.json +++ b/simpletuner/examples/kandinsky5-image-6b-i2i.lycoris-lokr/config.json @@ -12,7 +12,7 @@ "lora_rank": 128, "lora_type": "lycoris", "lr_scheduler": "constant", - "lycoris_config": "config/examples/kandinsky5-image.lycoris-lokr/lycoris_config.json", + "lycoris_config": "config/examples/kandinsky5-image-6b-i2i.lycoris-lokr/lycoris_config.json", "max_train_steps": 100, "model_family": "kandinsky5-image", "model_type": "lora", diff --git a/simpletuner/examples/kandinsky5-image-6b-t2i.lycoris-lokr/config.json b/simpletuner/examples/kandinsky5-image-6b-t2i.lycoris-lokr/config.json index 98338fb27..46cd15bc8 100644 --- a/simpletuner/examples/kandinsky5-image-6b-t2i.lycoris-lokr/config.json +++ b/simpletuner/examples/kandinsky5-image-6b-t2i.lycoris-lokr/config.json @@ -12,9 +12,9 @@ "lora_rank": 128, "lora_type": "lycoris", "lr_scheduler": "constant", - "lycoris_config": "config/examples/kandinsky5-image.lycoris-lokr/lycoris_config.json", + "lycoris_config": "config/examples/kandinsky5-image-6b-t2i.lycoris-lokr/lycoris_config.json", "max_train_steps": 100, - "model_family": "kandinsky5-image", + "model_family": "kandinsky5_image", "model_type": "lora", "num_eval_images": 25, "num_train_epochs": 0, diff --git a/simpletuner/examples/kandinsky5-video-2b-t2v.peft-lora/config.json b/simpletuner/examples/kandinsky5-video-2b-t2v.peft-lora/config.json index 3e1e52116..f6a58ab60 100644 --- a/simpletuner/examples/kandinsky5-video-2b-t2v.peft-lora/config.json +++ b/simpletuner/examples/kandinsky5-video-2b-t2v.peft-lora/config.json @@ -17,7 +17,7 @@ "max_train_steps": 200, "minimum_image_size": 0, "mixed_precision": "bf16", - "model_family": "kandinsky5-video", + "model_family": "kandinsky5_video", "model_type": "lora", "num_train_epochs": 0, "optimizer": "adamw_bf16", diff --git a/simpletuner/examples/kolors.peft-lora/config.json b/simpletuner/examples/kolors.peft-lora/config.json new file mode 100644 index 000000000..7b073cf3c --- /dev/null +++ b/simpletuner/examples/kolors.peft-lora/config.json @@ -0,0 +1,39 @@ +{ + "base_model_precision": "int8-quanto", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 50, + "checkpoints_total_limit": 5, + "data_backend_config": "config/examples/multidatabackend-small-dreambooth-512px.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-kolors-peft-lora", + "learning_rate": 1e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant", + "max_train_steps": 100, + "mixed_precision": "bf16", + "model_family": "kolors", + "model_flavour": "1.0", + "model_type": "lora", + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/kolors.peft-lora", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "quantize_via": "cpu", + "report_to": "none", + "resolution": 1024, + "resolution_type": "pixel_area", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "validation_guidance": 5.0, + "validation_num_inference_steps": 20, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "1024x1024", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/longcat-image.peft-lora/config.json b/simpletuner/examples/longcat-image.peft-lora/config.json index d3d70abb7..ebca4cb48 100644 --- a/simpletuner/examples/longcat-image.peft-lora/config.json +++ b/simpletuner/examples/longcat-image.peft-lora/config.json @@ -1,5 +1,6 @@ { "base_model_precision": "int8-quanto", + "attention_mechanism": "native-flash", "checkpoint_step_interval": 10, "data_backend_config": "config/examples/multidatabackend-small-dreambooth-512px.json", "disable_bucket_pruning": true, diff --git a/simpletuner/examples/longcat-video.peft-lora+ramtorch/config.json b/simpletuner/examples/longcat-video.peft-lora+ramtorch/config.json new file mode 100644 index 000000000..0988a6ec4 --- /dev/null +++ b/simpletuner/examples/longcat-video.peft-lora+ramtorch/config.json @@ -0,0 +1,50 @@ +{ + "attention_mechanism": "flash-attn-3-hub", + "base_model_precision": "no_change", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 100, + "checkpoints_total_limit": 5, + "compress_disk_cache": true, + "data_backend_config": "config/examples/multidatabackend-small-video-480p+81f.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-longcat-video-peft-lora-ramtorch", + "learning_rate": 6e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant_with_warmup", + "lr_warmup_steps": 50, + "max_grad_norm": 0.1, + "max_train_steps": 200, + "mixed_precision": "bf16", + "model_family": "longcat_video", + "model_flavour": "final", + "model_type": "lora", + "num_train_epochs": 0, + "offload_during_startup": true, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/longcat-video.peft-lora+ramtorch", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "ramtorch": true, + "ramtorch_target_modules": "blocks.*", + "ramtorch_transformer_percent": 100, + "report_to": "none", + "resolution": 480, + "resolution_type": "pixel_area", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "vae_enable_slicing": true, + "vae_enable_tiling": true, + "validation_guidance": 5.0, + "validation_negative_prompt": "blurry, cropped, ugly", + "validation_num_inference_steps": 30, + "validation_num_video_frames": 81, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "832x480", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/ltxvideo-0.9.5-t2v.peft-lora/config.json b/simpletuner/examples/ltxvideo-0.9.5-t2v.peft-lora/config.json new file mode 100644 index 000000000..54e319312 --- /dev/null +++ b/simpletuner/examples/ltxvideo-0.9.5-t2v.peft-lora/config.json @@ -0,0 +1,47 @@ +{ + "base_model_precision": "int8-torchao", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 100, + "checkpoints_total_limit": 5, + "compress_disk_cache": true, + "data_backend_config": "config/examples/multidatabackend-small-video-480p+2s24fps.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-ltxvideo-0-9-5-t2v-peft-lora", + "learning_rate": 6e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant_with_warmup", + "lr_warmup_steps": 50, + "max_grad_norm": 0.1, + "max_train_steps": 200, + "mixed_precision": "bf16", + "model_family": "ltxvideo", + "model_flavour": "0.9.5", + "pretrained_model_name_or_path": "Lightricks/LTX-Video-0.9.5", + "model_type": "lora", + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/ltxvideo-0.9.5-t2v.peft-lora", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "quantize_via": "cpu", + "report_to": "none", + "resolution": 480, + "resolution_type": "pixel_area", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "vae_enable_slicing": true, + "vae_enable_tiling": true, + "validation_guidance": 5.0, + "validation_negative_prompt": "blurry, cropped, ugly", + "validation_num_inference_steps": 40, + "validation_num_video_frames": 49, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "768x512", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard/config.json b/simpletuner/examples/ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard/config.json index 52d8b0387..a91198720 100644 --- a/simpletuner/examples/ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard/config.json +++ b/simpletuner/examples/ltxvideo2-2.3-dev-720p-single-gpu.peft-lora+sdnq-hadamard/config.json @@ -34,7 +34,7 @@ "resume_from_checkpoint": "latest", "sdnq_compile_mode": "compile", "sdnq_group_size": -1, - "sdnq_hadamard_group_size": 128, + "sdnq_hadamard_group_size": 256, "sdnq_use_hadamard": true, "sdnq_use_quantized_matmul": true, "seed": 42, diff --git a/simpletuner/examples/omnigen.lycoris-lokr/config.json b/simpletuner/examples/omnigen.lycoris-lokr/config.json index 1ba698f3b..0262635bf 100644 --- a/simpletuner/examples/omnigen.lycoris-lokr/config.json +++ b/simpletuner/examples/omnigen.lycoris-lokr/config.json @@ -56,5 +56,5 @@ "validation_prompt_library": false, "validation_resolution": "1024x1024", "validation_seed": 42, - "validation_steps": 10, + "validation_steps": 10 } diff --git a/simpletuner/examples/pixart.lycoris-lokr/config.json b/simpletuner/examples/pixart.lycoris-lokr/config.json index 5f984e5fe..02e01afe2 100644 --- a/simpletuner/examples/pixart.lycoris-lokr/config.json +++ b/simpletuner/examples/pixart.lycoris-lokr/config.json @@ -14,7 +14,7 @@ "lr_scheduler": "constant", "lycoris_config": "config/examples/pixart.lycoris-lokr/lycoris_config.json", "max_train_steps": 100, - "model_family": "pixart_sigma", + "model_family": "pixart", "model_type": "lora", "num_eval_images": 25, "num_train_epochs": 0, diff --git a/simpletuner/examples/sanavideo-2b-480p.peft-lora/config.json b/simpletuner/examples/sanavideo-2b-480p.peft-lora/config.json new file mode 100644 index 000000000..c22febd46 --- /dev/null +++ b/simpletuner/examples/sanavideo-2b-480p.peft-lora/config.json @@ -0,0 +1,47 @@ +{ + "base_model_precision": "int8-torchao", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 100, + "checkpoints_total_limit": 5, + "compress_disk_cache": true, + "data_backend_config": "config/examples/multidatabackend-small-video-480p+2s24fps.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-sanavideo-2b-480p-peft-lora", + "learning_rate": 6e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant_with_warmup", + "lr_warmup_steps": 50, + "max_grad_norm": 0.1, + "max_train_steps": 200, + "mixed_precision": "bf16", + "model_family": "sanavideo", + "model_flavour": "2b-480p", + "pretrained_model_name_or_path": "Efficient-Large-Model/SANA-Video_2B_480p_diffusers", + "model_type": "lora", + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/sanavideo-2b-480p.peft-lora", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "quantize_via": "cpu", + "report_to": "none", + "resolution": 480, + "resolution_type": "pixel_area", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "vae_enable_slicing": true, + "vae_enable_tiling": true, + "validation_guidance": 5.0, + "validation_negative_prompt": "blurry, cropped, ugly", + "validation_num_inference_steps": 30, + "validation_num_video_frames": 49, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "832x480", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/sd1x-dreamshaper.peft-lora/config.json b/simpletuner/examples/sd1x-dreamshaper.peft-lora/config.json new file mode 100644 index 000000000..3b748f87d --- /dev/null +++ b/simpletuner/examples/sd1x-dreamshaper.peft-lora/config.json @@ -0,0 +1,39 @@ +{ + "base_model_precision": "no_change", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 50, + "checkpoints_total_limit": 5, + "data_backend_config": "config/examples/multidatabackend-small-dreambooth-512px.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-sd1x-dreamshaper-peft-lora", + "learning_rate": 1e-4, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant", + "max_train_steps": 100, + "mixed_precision": "bf16", + "model_family": "sd1x", + "model_flavour": "dreamshaper", + "pretrained_model_name_or_path": "Lykon/dreamshaper-8", + "model_type": "lora", + "num_train_epochs": 0, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/sd1x-dreamshaper.peft-lora", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "report_to": "none", + "resolution": 512, + "resolution_type": "pixel", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "validation_guidance": 7.0, + "validation_num_inference_steps": 20, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "512x512", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/wan-s2v-14b-480p.peft-lora+ramtorch/config.json b/simpletuner/examples/wan-s2v-14b-480p.peft-lora+ramtorch/config.json new file mode 100644 index 000000000..34ef6b936 --- /dev/null +++ b/simpletuner/examples/wan-s2v-14b-480p.peft-lora+ramtorch/config.json @@ -0,0 +1,51 @@ +{ + "attention_mechanism": "flash-attn-3-hub", + "base_model_precision": "no_change", + "caption_dropout_probability": 0.1, + "checkpoint_step_interval": 100, + "checkpoints_total_limit": 5, + "compress_disk_cache": true, + "data_backend_config": "config/examples/multidatabackend-small-video-480p+81f.json", + "disable_bucket_pruning": true, + "gradient_checkpointing": true, + "hub_model_id": "simpletuner-example-wan-s2v-14b-480p-peft-lora-ramtorch", + "learning_rate": 6e-5, + "lora_rank": 16, + "lora_type": "standard", + "lr_scheduler": "constant_with_warmup", + "lr_warmup_steps": 50, + "max_grad_norm": 0.1, + "max_train_steps": 200, + "mixed_precision": "bf16", + "model_family": "wan_s2v", + "model_flavour": "s2v-14b-2.2", + "pretrained_model_name_or_path": "tolgacangoz/Wan2.2-S2V-14B-Diffusers", + "model_type": "lora", + "num_train_epochs": 0, + "offload_during_startup": true, + "optimizer": "adamw_bf16", + "output_dir": "output/examples/wan-s2v-14b-480p.peft-lora+ramtorch", + "push_checkpoints_to_hub": false, + "push_to_hub": false, + "ramtorch": true, + "ramtorch_target_modules": "blocks.*", + "ramtorch_transformer_percent": 100, + "report_to": "none", + "resolution": 480, + "resolution_type": "pixel_area", + "seed": 42, + "train_batch_size": 1, + "use_ema": false, + "vae_batch_size": 1, + "vae_enable_slicing": true, + "vae_enable_tiling": true, + "validation_guidance": 5.0, + "validation_negative_prompt": "blurry, cropped, ugly", + "validation_num_inference_steps": 30, + "validation_num_video_frames": 81, + "validation_prompt": "domokun holding a sign", + "validation_prompt_library": false, + "validation_resolution": "832x480", + "validation_seed": 42, + "validation_steps": 50 +} diff --git a/simpletuner/examples/wan2.1-t2v-14b-480p-8xh100.peft-lora+cp-fa3/config.json b/simpletuner/examples/wan2.1-t2v-14b-480p-8xh100.peft-lora+cp-fa3/config.json index ef56c0037..fe05d14f6 100644 --- a/simpletuner/examples/wan2.1-t2v-14b-480p-8xh100.peft-lora+cp-fa3/config.json +++ b/simpletuner/examples/wan2.1-t2v-14b-480p-8xh100.peft-lora+cp-fa3/config.json @@ -1,12 +1,12 @@ { "aspect_bucket_rounding": 2, "attention_mechanism": "flash-attn-3-hub", - "base_model_precision": "no_change", + "base_model_precision": "fp8-torchao", "caption_dropout_probability": 0.1, "checkpoint_step_interval": 100, "checkpoints_total_limit": 5, "compress_disk_cache": true, - "context_parallel_comm_strategy": "allgather", + "context_parallel_comm_strategy": "alltoall", "context_parallel_size": 2, "data_backend_config": "config/examples/multidatabackend-small-video-480p+81f.json", "delete_problematic_images": false, @@ -16,6 +16,7 @@ "grad_clip_method": "value", "gradient_accumulation_steps": 1, "gradient_checkpointing": true, + "gradient_checkpointing_interval": 1, "hub_model_id": "simpletuner-wan2.1-t2v-14b-480p-8xh100-lora-test", "strict_epoch_limit": false, "learning_rate": 6e-5, diff --git a/simpletuner/helpers/caching/vae.py b/simpletuner/helpers/caching/vae.py index 8f16af2c0..1fb57a4d4 100644 --- a/simpletuner/helpers/caching/vae.py +++ b/simpletuner/helpers/caching/vae.py @@ -400,9 +400,17 @@ def _handle_metadata_filtered_sample( delete_from_backend: bool = False, details: dict = None, ) -> None: + buckets = getattr(self.metadata_backend, "aspect_ratio_bucket_indices", None) + shard_lengths = ( + {shard: len(images) for shard, images in buckets.items()} + if isinstance(buckets, dict) and getattr(self.metadata_backend, "read_only", False) + else None + ) removed = self._remove_image_from_metadata_backend(self.metadata_backend, filepath, bucket) if removed: self._record_filter_stat(reason, bucket, filepath) + if shard_lengths is not None: + self._restore_split_shard_lengths(shard_lengths) self._queue_metadata_filter_action( filepath=filepath, bucket=bucket, @@ -411,6 +419,20 @@ def _handle_metadata_filtered_sample( details=details, ) + def _restore_split_shard_lengths(self, shard_lengths: dict) -> None: + """Refill a post-split shard the filter just shortened. + + Every data-parallel rank has to schedule the same number of samples, so removing an + entry from one rank's shard desynchronises them. The removed slots are refilled by + repeating the final surviving sample, which is how the splitter pads. A bucket the + filter emptied has nothing to repeat and stays empty until the next split. + """ + for bucket, length in shard_lengths.items(): + images = self.metadata_backend.aspect_ratio_bucket_indices.get(bucket) + if not images or len(images) >= length: + continue + images.extend([images[-1]] * (length - len(images))) + def _handle_nsfw_rejected_sample(self, filepath: str, bucket: str, classification: dict) -> None: rejected_entry = { "filepath": filepath, @@ -1034,6 +1056,7 @@ def prepare_video_latents(self, samples): "wan_s2v", "sanavideo", "kandinsky5-video", + "kandinsky5_video", "hunyuanvideo", "longcat_video", ]: @@ -1313,14 +1336,24 @@ def encode_images(self, images, filepaths, load_from_cache=True): if self.dataset_type_enum is DatasetType.AUDIO: debug_filepaths = [filepaths[i] for i in uncached_image_indices] + debug_latents = ( + latents_uncached["latents"] + if isinstance(latents_uncached, dict) and "latents" in latents_uncached + else latents_uncached + ) for idx, fp in enumerate(debug_filepaths): - self._log_audio_tensor_stats("audio_latents_processed", latents_uncached[idx], fp) + self._log_audio_tensor_stats("audio_latents_processed", debug_latents[idx], fp) latents_uncached = self.model.scale_vae_latents_for_cache(latents_uncached, self.vae) if self.dataset_type_enum is DatasetType.AUDIO: debug_filepaths = [filepaths[i] for i in uncached_image_indices] + debug_latents = ( + latents_uncached["latents"] + if isinstance(latents_uncached, dict) and "latents" in latents_uncached + else latents_uncached + ) for idx, fp in enumerate(debug_filepaths): - self._log_audio_tensor_stats("audio_latents_scaled", latents_uncached[idx], fp) + self._log_audio_tensor_stats("audio_latents_scaled", debug_latents[idx], fp) if isinstance(latents_uncached, dict) and "latents" in latents_uncached: raw_latents = latents_uncached["latents"] num_samples = raw_latents.shape[0] @@ -1604,6 +1637,12 @@ def _prepare_audio_sample(self, filepath: str, raw_sample): waveform, sample_rate = self._coerce_audio_waveform(raw_sample, metadata, filepath) if waveform is None: return None + waveform, metadata = self._truncate_audio_waveform_to_max_duration( + waveform=waveform, + sample_rate=sample_rate, + metadata=metadata, + filepath=filepath, + ) waveform, metadata = self._align_audio_waveform_to_video( waveform=waveform, sample_rate=sample_rate, @@ -1662,6 +1701,44 @@ def _align_audio_waveform_to_video(self, waveform, sample_rate, metadata: dict, metadata["duration_seconds"] = float(waveform.shape[-1]) / float(sample_rate) return waveform, metadata + def _truncate_audio_waveform_to_max_duration(self, waveform, sample_rate, metadata: dict, filepath: str): + if sample_rate is None: + return waveform, metadata + backend_config = StateTracker.get_data_backend_config(data_backend_id=self.id) or {} + audio_config = backend_config.get("audio") or {} + max_duration = audio_config.get("max_duration_seconds") + if max_duration is None: + return waveform, metadata + try: + max_duration = float(max_duration) + except (TypeError, ValueError): + return waveform, metadata + if max_duration <= 0: + return waveform, metadata + + target_samples = int(round(max_duration * float(sample_rate))) + if target_samples <= 0 or waveform.shape[-1] <= target_samples: + return waveform, metadata + + truncation_mode = (metadata.get("truncation_mode") or audio_config.get("truncation_mode") or "beginning").lower() + if truncation_mode == "end": + start = waveform.shape[-1] - target_samples + elif truncation_mode == "random": + start = random.randint(0, waveform.shape[-1] - target_samples) + else: + start = 0 + waveform = waveform[:, start : start + target_samples].contiguous() + metadata = dict(metadata) + metadata["num_samples"] = waveform.shape[-1] + metadata["duration_seconds"] = float(waveform.shape[-1]) / float(sample_rate) + metadata["truncated_duration_seconds"] = metadata["duration_seconds"] + logger.debug( + "Truncated audio sample %s to %.2fs for audio.max_duration_seconds.", + filepath, + metadata["duration_seconds"], + ) + return waveform, metadata + def _coerce_audio_waveform(self, sample, metadata: dict, filepath: str): waveform = None sample_rate = None diff --git a/simpletuner/helpers/configuration/env_file.py b/simpletuner/helpers/configuration/env_file.py index 5faa513cd..a88812a56 100644 --- a/simpletuner/helpers/configuration/env_file.py +++ b/simpletuner/helpers/configuration/env_file.py @@ -32,8 +32,13 @@ "MODEL_TYPE": "--model_type", "MODEL_NAME": "--pretrained_model_name_or_path", "MODEL_FAMILY": "--model_family", + "QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH": "--qwen_text_encoder_model_name_or_path", "TRAIN_BATCH_SIZE": "--train_batch_size", "USE_GRADIENT_CHECKPOINTING": "--gradient_checkpointing", + "GRADIENT_CHECKPOINTING_OFFLOAD_ATTENTION": "--gradient_checkpointing_offload_attention", + "GRADIENT_CHECKPOINTING_OFFLOAD_PREFETCH": "--gradient_checkpointing_offload_prefetch", + "GRADIENT_CHECKPOINTING_OFFLOAD_PIN_MEMORY_MAX_BUCKETS": "--gradient_checkpointing_offload_pin_memory_max_buckets", + "GRADIENT_CHECKPOINTING_SEGMENT_STRIDE": "--gradient_checkpointing_segment_stride", "ENABLE_CHUNKED_FEED_FORWARD": "--enable_chunked_feed_forward", "FEED_FORWARD_CHUNK_SIZE": "--feed_forward_chunk_size", "CAPTION_DROPOUT_PROBABILITY": "--caption_dropout_probability", diff --git a/simpletuner/helpers/data_backend/factory.py b/simpletuner/helpers/data_backend/factory.py index 6694ad272..7e9d6bbad 100644 --- a/simpletuner/helpers/data_backend/factory.py +++ b/simpletuner/helpers/data_backend/factory.py @@ -809,7 +809,7 @@ def _maybe_convert_pixel_edge(field_name: str) -> None: is_i2v_flavour = model_flavour.startswith("i2v") force_i2v = ( (model_family == "wan" and model_flavour.startswith("i2v-")) - or (model_family == "kandinsky5-video" and is_i2v_flavour) + or (model_family in ("kandinsky5-video", "kandinsky5_video") and is_i2v_flavour) or (model_family == "hunyuanvideo" and is_i2v_flavour) ) @@ -1907,7 +1907,7 @@ def _inject_i2v_conditioning_configs(self, data_backend_config: List[Dict[str, A matching conditioning-image-embed backends when none were supplied explicitly. """ model_family = str(getattr(self.args, "model_family", "") or "") - if model_family.lower() not in ["wan", "kandinsky5-video", "hunyuanvideo"]: + if model_family.lower() not in ["wan", "kandinsky5-video", "kandinsky5_video", "hunyuanvideo"]: return data_backend_config auto_embed_configs: List[Dict[str, Any]] = [] @@ -3283,7 +3283,7 @@ def _handle_bucket_operations( f" You have to reduce your batch size, or increase your dataset size (id={init_backend['id']})." ) - apply_padding = True if not self.args.max_train_steps or self.args.max_train_steps == 0 else False + apply_padding = not self.args.max_train_steps or self.args.allow_dataset_oversubscription if backend.get("auto_generated", False): # when we're duplicating a metadata set, it's already split between processes. diff --git a/simpletuner/helpers/metadata/backends/base.py b/simpletuner/helpers/metadata/backends/base.py index 0566774f1..f4b6f5fcb 100644 --- a/simpletuner/helpers/metadata/backends/base.py +++ b/simpletuner/helpers/metadata/backends/base.py @@ -778,10 +778,17 @@ def split_buckets_between_processes(self, gradient_accumulation_steps=1, apply_p if self.bucket_report: self.bucket_report.set_constraints(effective_batch_size=effective_batch_size) + backend_config = StateTracker.get_data_backend_config(self.id) or {} + configured_repeats = int(backend_config.get("repeats") or 0) + user_set_repeats = configured_repeats > 0 + auto_repeat_count = None + # Early validation: check if configuration is mathematically impossible buckets_that_will_fail = [] for bucket, images in self.aspect_ratio_bucket_indices.items(): total_img_count_incl_repeats = len(images) * (self.repeats + 1) + if not images: + continue if total_img_count_incl_repeats < effective_batch_size: buckets_that_will_fail.append( { @@ -799,25 +806,22 @@ def split_buckets_between_processes(self, gradient_accumulation_steps=1, apply_p needed_repeats = ceil(effective_batch_size / images) - 1 min_repeats_needed[bucket_info["bucket"]] = needed_repeats + # The documented repeat setting applies to the whole backend, so + # use the maximum requirement and apply it consistently to every + # bucket rather than assigning a different repeat count per bucket. max_needed_repeats = max(min_repeats_needed.values()) + allow_oversubscription = StateTracker.get_args().allow_dataset_oversubscription # Check if dataset oversubscription is allowed - args = StateTracker.get_args() - allow_oversubscription = args.allow_dataset_oversubscription - - # Check if user manually configured repeats in their backend config - backend_config = StateTracker.get_data_backend_config(self.id) or {} - user_set_repeats = "repeats" in backend_config - if allow_oversubscription and not user_set_repeats: # Automatically adjust repeats to make training possible original_repeats = self.repeats - self.repeats = max_needed_repeats + auto_repeat_count = max_needed_repeats logger.warning( - f"(id={self.id}) Dataset oversubscription enabled: automatically increasing repeats from {original_repeats} to {self.repeats}\n" + f"(id={self.id}) Dataset oversubscription enabled: automatically increasing repeats from {original_repeats} to {auto_repeat_count}\n" f" - This allows training with {total_samples} samples across {num_processes} GPUs\n" f" - Effective batch size: {effective_batch_size}\n" - f" - Each sample will be seen {self.repeats + 1} times per epoch" + f" - Logical repeat factor before per-bucket batch padding: {auto_repeat_count + 1}" ) # Validation passed with adjustment, continue else: @@ -874,10 +878,27 @@ def split_buckets_between_processes(self, gradient_accumulation_steps=1, apply_p shuffle_seed = broadcast_object_from_main(shuffle_seed if self.accelerator.is_main_process else None) for bucket, images in self.aspect_ratio_bucket_indices.items(): + if not images: + new_aspect_ratio_bucket_indices[bucket] = [] + continue if should_shuffle_contents: logger.debug(f"Shuffling bucket {bucket} contents.") images = images.copy() random.Random(f"{shuffle_seed}:{self.id}:{bucket}").shuffle(images) + + if auto_repeat_count is not None: + logical_count = len(images) * (auto_repeat_count + 1) + scheduled_count = ceil(logical_count / effective_batch_size) * effective_batch_size + local_count = scheduled_count // effective_dp_size + start_idx = dp_rank * local_count + images_split = [images[(start_idx + offset) % len(images)] for offset in range(local_count)] + logger.debug( + f"(id={self.id}) Bucket {bucket}: logical samples={logical_count}, " + f"scheduled samples={scheduled_count}, local samples={local_count}" + ) + new_aspect_ratio_bucket_indices[bucket] = images_split + continue + total_img_count_incl_repeats = len(images) * (self.repeats + 1) num_batches = ceil(total_img_count_incl_repeats / effective_batch_size) trim_limit = num_batches * effective_batch_size @@ -903,33 +924,21 @@ def split_buckets_between_processes(self, gradient_accumulation_steps=1, apply_p effective_batch_size=effective_batch_size, ) - # Split data by DP rank (not global rank) when context parallelism is enabled. - # This ensures all ranks in a CP group get the same data shard. - if cp_size > 1: - # Custom CP-aware splitting: split by dp_rank instead of process_index - chunk_size = ceil(len(trimmed_images) / effective_dp_size) if effective_dp_size > 0 else len(trimmed_images) - start_idx = dp_rank * chunk_size - end_idx = min(start_idx + chunk_size, len(trimmed_images)) - images_split = trimmed_images[start_idx:end_idx] - - # Handle padding if requested (for uniform batch sizes) - if apply_padding and len(images_split) < chunk_size and len(images_split) > 0: - padding_needed = chunk_size - len(images_split) - # Pad by repeating elements from the split - padding = (images_split * ((padding_needed // len(images_split)) + 1))[:padding_needed] - images_split = images_split + padding - - new_aspect_ratio_bucket_indices[bucket] = images_split - else: - # Standard splitting when CP is disabled - with self.accelerator.split_between_processes(trimmed_images, apply_padding=apply_padding) as images_split: - new_aspect_ratio_bucket_indices[bucket] = images_split + samples_per_rank, extra_samples = divmod(len(trimmed_images), effective_dp_size) + start_idx = dp_rank * samples_per_rank + min(dp_rank, extra_samples) + local_size = samples_per_rank + int(dp_rank < extra_samples) + images_split = trimmed_images[start_idx : start_idx + local_size] + if apply_padding: + target_size = samples_per_rank + int(extra_samples > 0) + if trimmed_images and len(images_split) < target_size: + images_split += [trimmed_images[-1]] * (target_size - len(images_split)) + new_aspect_ratio_bucket_indices[bucket] = images_split self.aspect_ratio_bucket_indices = new_aspect_ratio_bucket_indices post_total = sum([len(bucket) for bucket in self.aspect_ratio_bucket_indices.values()]) if self.bucket_report: self.bucket_report.record_bucket_snapshot("post_split", self.aspect_ratio_bucket_indices) - if total_samples != post_total: + if self.accelerator.num_processes > 1 or total_samples != post_total: self.read_only = True # Check if this backend has no samples after splitting (can happen with multi-GPU setups) @@ -944,14 +953,28 @@ def split_buckets_between_processes(self, gradient_accumulation_steps=1, apply_p logger.debug(f"Count of items after split: {post_total}") + def seen_occurrence_count(self, image_path) -> int: + """Return consumed occurrences, accepting boolean values from older checkpoints.""" + value = self.seen_images.get(image_path, 0) + if isinstance(value, bool): + return int(value) + if not isinstance(value, int): + raise TypeError(f"Invalid seen occurrence count for {image_path!r}: {value!r}") + if value < 0: + raise ValueError(f"Seen occurrence count cannot be negative for {image_path!r}: {value}") + return int(value) + def mark_as_seen(self, image_path): - self.seen_images[image_path] = True + self.seen_images[image_path] = self.seen_occurrence_count(image_path) + 1 def mark_batch_as_seen(self, image_paths): - self.seen_images.update({image_path: True for image_path in image_paths}) + for image_path in image_paths: + self.mark_as_seen(image_path) - def is_seen(self, image_path): - return self.seen_images.get(image_path, False) + def is_seen(self, image_path, occurrence_index: int = 0): + if isinstance(self.seen_images.get(image_path, 0), bool): + return self.seen_images[image_path] + return self.seen_occurrence_count(image_path) > occurrence_index def reset_seen_images(self): self.seen_images.clear() diff --git a/simpletuner/helpers/metadata/backends/huggingface.py b/simpletuner/helpers/metadata/backends/huggingface.py index 472550373..5d8da3f53 100644 --- a/simpletuner/helpers/metadata/backends/huggingface.py +++ b/simpletuner/helpers/metadata/backends/huggingface.py @@ -165,6 +165,7 @@ def _extract_captions_to_dict(self) -> Dict[str, Union[str, List[str]]]: except Exception as e: logger.warning(f"Error loading caption cache, will regenerate: {e}") logger.info("Extracting captions from Hugging Face dataset...") + indices = self._limited_dataset_indices() def process_item(idx): item = self.data_backend.dataset[idx] @@ -179,13 +180,12 @@ def process_item(idx): return None captions = {} - total_items = len(self.data_backend.dataset) with concurrent.futures.ThreadPoolExecutor() as executor: results = list( tqdm( - executor.map(process_item, range(total_items)), + executor.map(process_item, indices), desc="Extracting captions", - total=total_items, + total=len(indices), ncols=100, mininterval=0.5, ascii=True, @@ -206,6 +206,18 @@ def process_item(idx): logger.info(f"Extracted {len(captions)} captions from dataset") return captions + def _limited_dataset_indices(self) -> List[int]: + total_items = len(self.data_backend.dataset) + all_files = [f"{idx}.{self.file_extension}" for idx in range(total_items)] + limited_files = self._apply_max_num_samples_limit(all_files) + indices = [] + for virtual_path in limited_files: + try: + indices.append(int(str(virtual_path).rsplit(".", 1)[0])) + except (TypeError, ValueError): + continue + return indices + def _extract_caption_from_item(self, item: Dict) -> Optional[Union[str, List[str]]]: caption = None # handle list of caption columns @@ -758,23 +770,25 @@ def compute_aspect_ratio_bucket_indices(self, ignore_existing_cache: bool = Fals aspect_ratio_bucket_updates = {} metadata_updates = {} - total_items = len(self.data_backend.dataset) + indices = self._limited_dataset_indices() if self.bucket_report: - pending_items = max(total_items - len(existing_files), 0) + pending_items = max(len(indices) - len(existing_files), 0) self.bucket_report.record_stage( "new_files_to_process", sample_count=pending_items, ignore_existing_cache=ignore_existing_cache, ) - for idx in tqdm( - range(total_items), - desc="Processing HF dataset items", - total=total_items, - leave=False, - ncols=100, + for progress_index, idx in enumerate( + tqdm( + indices, + desc="Processing HF dataset items", + total=len(indices), + leave=False, + ncols=100, + ) ): if progress_callback is not None: - progress_callback(idx + 1, total_items) + progress_callback(progress_index + 1, len(indices)) virtual_path = f"{idx}.{self.file_extension}" diff --git a/simpletuner/helpers/metadata/backends/parquet.py b/simpletuner/helpers/metadata/backends/parquet.py index 4c7f88774..fc918ee72 100644 --- a/simpletuner/helpers/metadata/backends/parquet.py +++ b/simpletuner/helpers/metadata/backends/parquet.py @@ -280,6 +280,10 @@ def reload_cache(self, set_config: bool = True): ) # Load filtering statistics if present self.filtering_statistics = cache_data.get("filtering_statistics") + try: + self.load_image_metadata() + except Exception as e: + logger.warning(f"Error loading parquet metadata cache, continuing with empty metadata: {e}") else: logger.warning("No cache file found, starting a fresh one.") @@ -344,9 +348,15 @@ def compute_aspect_ratio_bucket_indices(self, ignore_existing_cache: bool = Fals ignore_existing_cache=ignore_existing_cache, ) if not new_files: + if not self.image_metadata_loaded: + try: + self.load_image_metadata() + except Exception as e: + raise Exception(f"Error loading image metadata. Consider removing the metadata file manually: {e}") if self.bucket_report: self.bucket_report.update_statistics(statistics) self.bucket_report.record_bucket_snapshot("post_refresh", self.aspect_ratio_bucket_indices) + self._ensure_metadata_for_cached_buckets() return if ignore_existing_cache: @@ -421,14 +431,17 @@ def _get_first_value(self, series_or_scalar): if isinstance(series_or_scalar, pd.Series): series_or_scalar = series_or_scalar.iloc[0] # Just unwrap the first value elif isinstance(series_or_scalar, str): - series_or_scalar = float(series_or_scalar) if "." in series_or_scalar else int(series_or_scalar) + try: + series_or_scalar = float(series_or_scalar) if "." in series_or_scalar else int(series_or_scalar) + except ValueError: + return series_or_scalar # After unwrapping, if it's an np.int* or np.float*, cast to python int/float if isinstance(series_or_scalar, np.integer): return int(series_or_scalar) elif isinstance(series_or_scalar, np.floating): return float(series_or_scalar) - elif isinstance(series_or_scalar, (int, float)): + elif isinstance(series_or_scalar, (int, float, str, list, tuple, dict)): return series_or_scalar elif series_or_scalar is None: return None @@ -708,6 +721,11 @@ def _process_audio_bucket( duration_seconds = None overrides = {} + caption_column = self.parquet_config.get("caption_column") + caption_value = self._extract_audio_value(database_row, caption_column) + if caption_value: + overrides["tags"] = caption_value + overrides["prompt"] = caption_value lyrics_column = self.parquet_config.get("lyrics_column") if lyrics_column: lyrics_value = self._extract_audio_value(database_row, lyrics_column) diff --git a/simpletuner/helpers/metadata/utils/duplicator.py b/simpletuner/helpers/metadata/utils/duplicator.py index e7d06fd73..b11816f25 100644 --- a/simpletuner/helpers/metadata/utils/duplicator.py +++ b/simpletuner/helpers/metadata/utils/duplicator.py @@ -21,19 +21,7 @@ def copy_metadata(source_backend, target_backend): if source_meta is None or target_meta is None: raise ValueError(f"Both backends must have metadata_backend defined. Received {source_meta} \n\n {target_meta}") - logger.debug("Reloading metadata caches...") - prior_source_buckets = { - bucket: list(paths) for bucket, paths in getattr(source_meta, "aspect_ratio_bucket_indices", {}).items() - } - prior_source_metadata = dict(source_meta.image_metadata) if getattr(source_meta, "image_metadata", None) else None - source_meta.reload_cache(set_config=False) - if not source_meta.aspect_ratio_bucket_indices and prior_source_buckets: - logger.warning( - "Source metadata cache reload returned no buckets; restoring in-memory snapshot to avoid data loss." - ) - source_meta.aspect_ratio_bucket_indices = {bucket: list(paths) for bucket, paths in prior_source_buckets.items()} - if prior_source_metadata is not None: - source_meta.image_metadata = prior_source_metadata + logger.debug("Reloading target metadata cache...") target_meta.reload_cache(set_config=False) # Get the instance directories for path translation @@ -88,7 +76,7 @@ def copy_metadata(source_backend, target_backend): else: # Regular copy without path translation logger.info("Copying metadata without path translation") - target_meta.set_metadata(metadata_backend=source_meta, update_json=True) + target_meta.set_metadata(metadata_backend=source_meta, update_json=False) source_config = source_backend.get("config", {}) or {} target_config = target_backend.get("config", {}) or {} @@ -131,10 +119,8 @@ def copy_metadata(source_backend, target_backend): target_meta.config = target_config - # Save the updated metadata - target_meta.save_cache() - - # IMPORTANT: save_cache() only saves bucket indices, we must also save the image_metadata + # Bucket indices may be rank-local here; do not overwrite the canonical target cache. + target_meta.set_readonly() if hasattr(target_meta, "save_image_metadata"): target_meta.save_image_metadata() logger.debug("Saved image_metadata to disk") @@ -144,7 +130,6 @@ def copy_metadata(source_backend, target_backend): logger.info("Metadata copied successfully.") source_meta.print_debug_info() target_meta.print_debug_info() - target_meta.set_readonly() @staticmethod def generate_conditioning_datasets(global_config, source_backend_config): diff --git a/simpletuner/helpers/models/ace_step/model.py b/simpletuner/helpers/models/ace_step/model.py index cae5dd1fe..564faf210 100644 --- a/simpletuner/helpers/models/ace_step/model.py +++ b/simpletuner/helpers/models/ace_step/model.py @@ -244,7 +244,8 @@ def get_acceleration_presets(cls) -> list[AccelerationPreset]: def __init__(self, config: dict, accelerator): super().__init__(config, accelerator) - self.text_tokenizer_max_length = getattr(self.config, "tokenizer_max_length", 256) + default_text_length = 77 if str(getattr(self.config, "model_flavour", "")).startswith("v15") else 256 + self.text_tokenizer_max_length = int(getattr(self.config, "tokenizer_max_length", None) or default_text_length) self.mert_model = None self.hubert_model = None self.resampler_mert = None @@ -261,6 +262,7 @@ def __init__(self, config: dict, accelerator): self._v15_layout: Optional[Dict[str, str]] = None self._v15_layout_probe_base: Optional[str] = None self.silence_latent: Optional[torch.Tensor] = None + self.v15_condition_model_dtype = getattr(self.config, "weight_dtype", torch.float32) def get_lora_target_layers(self): manual_targets = self._get_peft_lora_target_modules() @@ -378,9 +380,13 @@ def _resolve_v15_layout(self, base_path: Optional[str] = None) -> Optional[Dict[ if available_variants: variant_dir = available_variants[0] + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") tokenizer_dir = shared_root / self.V15_SHARED_TEXT_ENCODER_SUBFOLDER vae_dir = shared_root / self.V15_SHARED_VAE_SUBFOLDER - if variant_dir is None or not tokenizer_dir.is_dir() or not vae_dir.is_dir(): + if variant_dir is None or not vae_dir.is_dir(): + self._v15_layout = None + return None + if not qwen_text_encoder_path and not tokenizer_dir.is_dir(): self._v15_layout = None return None @@ -404,7 +410,7 @@ def _resolve_v15_layout(self, base_path: Optional[str] = None) -> Optional[Dict[ self._v15_layout = { "root_path": str(shared_root), "variant_path": str(variant_dir), - "tokenizer_path": str(tokenizer_dir), + "tokenizer_path": qwen_text_encoder_path or str(tokenizer_dir), "vae_path": str(vae_dir), "silence_latent_path": str(silence_path), } @@ -434,6 +440,27 @@ def _trust_remote_code_enabled(self) -> bool: def _load_v15_hf_component(self, loader_cls, *, component_name: str, pretrained_model_name_or_path: str, **kwargs): trust_remote_code = self._trust_remote_code_enabled() + force_cpu_init_without_meta = bool(kwargs.pop("force_cpu_init_without_meta", False)) + original_get_init_context = None + if force_cpu_init_without_meta: + from transformers.modeling_utils import PreTrainedModel + + original_get_init_context = PreTrainedModel.get_init_context + original_get_init_context_func = original_get_init_context.__func__ + + def get_cpu_init_context(cls, dtype, is_quantized, _is_ds_init_called, allow_all_kernels): + contexts = original_get_init_context_func( + cls, + dtype, + is_quantized, + _is_ds_init_called, + allow_all_kernels, + ) + return [ + context for context in contexts if not (isinstance(context, torch.device) and context.type == "meta") + ] + + PreTrainedModel.get_init_context = classmethod(get_cpu_init_context) try: return loader_cls.from_pretrained( pretrained_model_name_or_path=pretrained_model_name_or_path, @@ -448,6 +475,9 @@ def _load_v15_hf_component(self, loader_cls, *, component_name: str, pretrained_ "`trust_remote_code=true` in the SimpleTuner configuration and retry." ) from exc raise + finally: + if original_get_init_context is not None: + PreTrainedModel.get_init_context = original_get_init_context def _build_v15_text_prompt(self, prompt: str, prompt_context: Optional[dict]) -> str: metadata = prompt_context or {} @@ -476,6 +506,18 @@ def _get_v15_text_embedding_layer(self): raise AttributeError("ACE-Step v1.5 text encoder does not expose an input embedding layer.") return embedding_layer + def _ensure_v15_condition_model_dtype(self): + model = getattr(self, "model", None) + if not self._is_v15_layout_active() or model is None: + return + component = self.unwrap_model(model=model) + try: + first_param = next(component.parameters()) + except StopIteration: + return + if first_param.dtype != self.v15_condition_model_dtype: + component.to(device=self.accelerator.device, dtype=self.v15_condition_model_dtype) + def _get_v15_silence_latent_slice(self, length: int, device, dtype) -> torch.Tensor: if self.silence_latent is None: raise ValueError("ACE-Step v1.5 requires silence_latent.pt to be loaded before preparing batches.") @@ -540,15 +582,79 @@ def _run_v15_encoder( ) refer_audio_order_mask = torch.zeros(text_hidden_states.shape[0], device=text_hidden_states.device, dtype=torch.long) full_model.encoder.eval() + model_config = getattr(full_model, "config", None) + original_attn_implementation = getattr(model_config, "_attn_implementation", None) with torch.no_grad(): - return full_model.encoder( - text_hidden_states=text_hidden_states, - text_attention_mask=text_attention_mask, - lyric_hidden_states=lyric_hidden_states, - lyric_attention_mask=lyric_attention_mask, - refer_audio_acoustic_hidden_states_packed=refer_audio_hidden, - refer_audio_order_mask=refer_audio_order_mask, + try: + if model_config is not None: + model_config._attn_implementation = "eager" + return full_model.encoder( + text_hidden_states=text_hidden_states, + text_attention_mask=text_attention_mask, + lyric_hidden_states=lyric_hidden_states, + lyric_attention_mask=lyric_attention_mask, + refer_audio_acoustic_hidden_states_packed=refer_audio_hidden, + refer_audio_order_mask=refer_audio_order_mask, + ) + finally: + if model_config is not None: + model_config._attn_implementation = original_attn_implementation + + def _configure_v15_flash_attention(self) -> None: + if not torch.cuda.is_available() or not hasattr(self.model, "config"): + return + + attn_implementation = "kernels-community/flash-attn2" + try: + import transformers.modeling_flash_attention_utils as flash_utils + from huggingface_hub import constants as hf_hub_constants + from transformers.integrations import hub_kernels + from transformers.integrations.flash_attention import flash_attention_forward + from transformers.masking_utils import ALL_MASK_ATTENTION_FUNCTIONS + from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS + + original_disable_telemetry = hf_hub_constants.HF_HUB_DISABLE_TELEMETRY + try: + # huggingface_hub 1.26 can emit an illegal trailing-semicolon + # User-Agent when telemetry is disabled and `kernels` downloads + # hub kernels. Keep the environment unchanged and only patch the + # cached constant for this direct kernel load. + hf_hub_constants.HF_HUB_DISABLE_TELEMETRY = False + kernel = hub_kernels.get_kernel_hub( + attn_implementation, + version=1, + trust_remote_code=True, + ) + finally: + hf_hub_constants.HF_HUB_DISABLE_TELEMETRY = original_disable_telemetry + flash_utils._loaded_implementation = attn_implementation + flash_utils._flash_fn = getattr(kernel, "flash_attn_func", None) + flash_utils._flash_varlen_fn = getattr(kernel, "flash_attn_varlen_func", None) + flash_utils._flash_with_kvcache_fn = getattr(kernel, "flash_attn_with_kvcache", None) + flash_utils._pad_fn = flash_utils._pad_input + flash_utils._unpad_fn = flash_utils._unpad_input + flash_utils._process_flash_kwargs_fn = flash_utils._lazy_define_process_function(flash_utils._flash_varlen_fn) + ALL_ATTENTION_FUNCTIONS.register(attn_implementation, flash_attention_forward) + ALL_MASK_ATTENTION_FUNCTIONS.register( + attn_implementation, + ALL_MASK_ATTENTION_FUNCTIONS["flash_attention_2"], ) + self.model.config._attn_implementation_internal = attn_implementation + except Exception as exc: + logger.warning( + "Could not register ACE-Step v1.5 FlashAttention 2 fallback kernel; falling back to eager attention: %s", + exc, + ) + self.model.config._attn_implementation = "eager" + + def _v15_flash_attention_mask(self, mask: Optional[torch.Tensor]) -> Optional[torch.Tensor]: + if mask is None: + return None + full_model = self.unwrap_model(model=self.model) + attn_implementation = getattr(getattr(full_model, "config", None), "_attn_implementation", None) + if attn_implementation != "eager" and torch.is_tensor(mask): + return None + return mask def _sample_v15_timesteps(self, batch_size: int, device, dtype) -> torch.Tensor: full_model = self.unwrap_model(model=self.model) @@ -643,6 +749,8 @@ def _resolve_checkpoint_base(self) -> str: [ f"{v15_variant_subdir}/*", f"{v15_variant_subdir}/{self.V15_SILENCE_LATENT_FILENAME}", + "acestep-v15-*/*", + "acestep-v15-*/silence_latent.pt", ] ) else: @@ -765,10 +873,14 @@ def load_model(self, move_to_device: bool = True): AutoModel, component_name="condition model", pretrained_model_name_or_path=v15_layout["variant_path"], - torch_dtype=self.config.weight_dtype, + torch_dtype=self.v15_condition_model_dtype, + low_cpu_mem_usage=False, + allow_all_kernels=True, + force_cpu_init_without_meta=True, ) if move_to_device: - self.model.to(self.accelerator.device, dtype=self.config.weight_dtype) + self.model.to(self.accelerator.device, dtype=self.v15_condition_model_dtype) + self._configure_v15_flash_attention() for module_name in ("encoder", "tokenizer", "detokenizer"): module = getattr(self.model, module_name, None) if module is None: @@ -1398,7 +1510,7 @@ def _prepare_conditioning_features(self, batch: dict): def get_trained_component(self, base_model: bool = False, unwrap_model: bool = True): if not self._is_v15_layout_active(): return super().get_trained_component(base_model=base_model, unwrap_model=unwrap_model) - component = self.model if base_model else getattr(self.model, "decoder", None) + component = self.model if unwrap_model: return self.unwrap_model(model=component) return component @@ -1444,11 +1556,11 @@ def add_lora_adapter(self): if self._ramtorch_enabled(): ramtorch_utils.register_lora_custom_module(self.lora_config) - self.model.decoder = get_peft_model(self.model.decoder, self.lora_config) + self.model = get_peft_model(self.model, self.lora_config) if getattr(self.config, "init_lora", None): addkeys, misskeys = load_lora_weights( - {self.MODEL_TYPE.value: self.model.decoder}, + {self.MODEL_TYPE.value: self.model}, self.config.init_lora, use_dora=getattr(self.config, "use_dora", False), ) @@ -1488,13 +1600,15 @@ def prepare_batch(self, batch: dict, state: dict) -> dict: batch["encoder_attention_mask"] = torch.ones(mask_shape, dtype=torch.long) device = self.accelerator.device - dtype = getattr(self.config, "weight_dtype", torch.float32) + dtype = self.v15_condition_model_dtype latents = latent_batch.to(device=device, dtype=dtype) batch_size = latents.shape[0] - normalized_lyrics = self._normalize_v15_lyrics_inputs(lyrics, batch_size=batch_size) attention_mask = batch["latent_attention_mask"].to(device=device, dtype=dtype) text_hidden_states = batch["prompt_embeds"].to(device=device, dtype=dtype) text_attention_mask = batch["encoder_attention_mask"].to(device=device, dtype=torch.long) + + self._ensure_v15_condition_model_dtype() + normalized_lyrics = self._normalize_v15_lyrics_inputs(lyrics, batch_size=batch_size) lyric_hidden_states, lyric_attention_mask = self._embed_v15_lyrics_batch(normalized_lyrics) encoder_hidden_states, encoder_attention_mask = self._run_v15_encoder( text_hidden_states=text_hidden_states, @@ -1523,6 +1637,8 @@ def prepare_batch(self, batch: dict, state: dict) -> dict: null_condition_emb = null_condition_emb.to(device=device, dtype=dtype).expand_as(encoder_hidden_states) encoder_hidden_states = torch.where(keep_mask > 0, encoder_hidden_states, null_condition_emb) + context_latents = self._build_v15_context_latents(latents.shape[1], latents.shape[0], device, dtype) + return { "latents": latents, "noise": noise, @@ -1531,7 +1647,7 @@ def prepare_batch(self, batch: dict, state: dict) -> dict: "attention_mask": attention_mask, "encoder_hidden_states": encoder_hidden_states, "encoder_attention_mask": encoder_attention_mask, - "context_latents": self._build_v15_context_latents(latents.shape[1], latents.shape[0], device, dtype), + "context_latents": context_latents, "flow_target": noise - latents, } @@ -1647,31 +1763,63 @@ def model_predict(self, prepared_batch: dict) -> Dict[str, object]: raise ValueError("ACE-Step transformer has not been loaded before model_predict was invoked.") if self._is_v15_layout_active(): + self._ensure_v15_condition_model_dtype() + full_model = self.unwrap_model(model=self.model) + decoder = getattr(full_model, "decoder", None) + if decoder is None and hasattr(full_model, "base_model"): + base_model = getattr(full_model.base_model, "model", full_model.base_model) + decoder = getattr(base_model, "decoder", None) + if decoder is None: + raise ValueError("ACE-Step v1.5 condition model does not expose a decoder module.") call_kwargs = { "hidden_states": prepared_batch["noisy_latents"], "timestep": prepared_batch["timesteps"], - "attention_mask": prepared_batch["attention_mask"], + "attention_mask": self._v15_flash_attention_mask(prepared_batch["attention_mask"]), "encoder_hidden_states": prepared_batch["encoder_hidden_states"], - "encoder_attention_mask": prepared_batch.get("encoder_attention_mask"), + "encoder_attention_mask": self._v15_flash_attention_mask(prepared_batch.get("encoder_attention_mask")), "context_latents": prepared_batch["context_latents"], "timestep_r": prepared_batch["timesteps"], } call_kwargs.update( self._get_flowmap_r_timestep_forward_kwargs( prepared_batch, - target=getattr(transformer, "forward", transformer), + target=getattr(decoder, "forward", decoder), kwarg_name="timestep_r", ) ) - output = transformer(**call_kwargs) + output = decoder(**call_kwargs) flow_pred = output[0] if isinstance(output, (tuple, list)) else getattr(output, "sample", None) if flow_pred is None and hasattr(output, "__getitem__"): flow_pred = output[0] if flow_pred is None: raise ValueError("ACE-Step v1.5 decoder did not return a flow prediction tensor.") if not torch.isfinite(flow_pred).all(): + + def _tensor_summary(name: str, tensor: Optional[torch.Tensor]) -> str: + if not torch.is_tensor(tensor): + return f"{name}=None" + finite = torch.isfinite(tensor) + if not finite.any(): + return f"{name}=shape{tuple(tensor.shape)} dtype={tensor.dtype} all_nonfinite" + sample = tensor.detach()[finite].float() + return ( + f"{name}=shape{tuple(tensor.shape)} dtype={tensor.dtype} " + f"min={sample.min().item()} max={sample.max().item()} mean={sample.mean().item()}" + ) + raise ValueError( - f"Non-finite model_prediction detected (min={flow_pred.min().item()}, max={flow_pred.max().item()})" + "Non-finite model_prediction detected: " + + "; ".join( + [ + _tensor_summary("noisy_latents", prepared_batch.get("noisy_latents")), + _tensor_summary("timesteps", prepared_batch.get("timesteps")), + _tensor_summary("attention_mask", prepared_batch.get("attention_mask")), + _tensor_summary("encoder_hidden_states", prepared_batch.get("encoder_hidden_states")), + _tensor_summary("encoder_attention_mask", prepared_batch.get("encoder_attention_mask")), + _tensor_summary("context_latents", prepared_batch.get("context_latents")), + _tensor_summary("flow_pred", flow_pred), + ] + ) ) return { "model_prediction": flow_pred, @@ -1736,6 +1884,11 @@ def loss(self, prepared_batch: dict, model_output, apply_conditioning_mask: bool """ import torch.nn.functional as F + if self._is_v15_layout_active(): + diffusion_loss = model_output.get("diffusion_loss") if isinstance(model_output, dict) else None + if diffusion_loss is not None: + return diffusion_loss + model_pred = model_output.get("model_prediction") if model_pred is None: model_pred = model_output.get("sample") @@ -1770,6 +1923,12 @@ def loss(self, prepared_batch: dict, model_output, apply_conditioning_mask: bool return loss def auxiliary_loss(self, model_output, prepared_batch: dict, loss: torch.Tensor, **kwargs): + if ( + self._is_v15_layout_active() + and isinstance(model_output, dict) + and model_output.get("diffusion_loss") is not None + ): + return loss, None loss, base_logs = super().auxiliary_loss(model_output=model_output, prepared_batch=prepared_batch, loss=loss) proj_losses = model_output.get("proj_losses") if not proj_losses: diff --git a/simpletuner/helpers/models/ace_step/transformer.py b/simpletuner/helpers/models/ace_step/transformer.py index 2286c9536..8d5c47612 100644 --- a/simpletuner/helpers/models/ace_step/transformer.py +++ b/simpletuner/helpers/models/ace_step/transformer.py @@ -37,6 +37,7 @@ set_flowmap_gate, validate_flowmap_deltatime_type, ) +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from .attention import LinearTransformerBlock, t2i_modulate from .lyrics_utils.lyric_encoder import ConformerEncoder as LyricEncoder @@ -364,6 +365,8 @@ def __init__( self.final_layer = T2IFinalLayer(self.inner_dim, patch_size=patch_size, out_channels=out_channels) self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" self._mps_fp32 = False self._logged_dtype_mismatch = False @@ -371,6 +374,12 @@ def __init__( def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.flowmap_deltatime_type = validate_flowmap_deltatime_type(deltatime_type, model_name="ACEStep") if self.delta_timestep_embedder is None: @@ -619,9 +628,14 @@ def decode( capture_idx = 0 for index_block, block in enumerate(self.transformer_blocks): - if self.training and self.gradient_checkpointing: + if self.training and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/anima/model.py b/simpletuner/helpers/models/anima/model.py index 3cbf8544d..80aff1c62 100644 --- a/simpletuner/helpers/models/anima/model.py +++ b/simpletuner/helpers/models/anima/model.py @@ -283,20 +283,21 @@ def load_model(self, move_to_device: bool = True): return super().load_model(move_to_device=move_to_device) def _prompt_tokenizer_sources(self) -> tuple[str, str]: + qwen_tokenizer_source = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") model_path = getattr(self.config, "pretrained_model_name_or_path", None) if isinstance(model_path, str): model_dir = Path(model_path) qwen_dir = model_dir / "tokenizer" t5_dir = model_dir / "t5_tokenizer" if qwen_dir.is_dir() and t5_dir.is_dir(): - return str(qwen_dir), str(t5_dir) + return qwen_tokenizer_source or str(qwen_dir), str(t5_dir) qwen_dir = model_dir / "prompt_tokenizer_qwen" t5_dir = model_dir / "prompt_tokenizer_t5" if qwen_dir.is_dir() and t5_dir.is_dir(): - return str(qwen_dir), str(t5_dir) + return qwen_tokenizer_source or str(qwen_dir), str(t5_dir) if self._uses_diffusers_repo_layout(): - return f"{model_path}::tokenizer", f"{model_path}::t5_tokenizer" - return _QWEN_TOKENIZER_SOURCE, _T5_TOKENIZER_SOURCE + return qwen_tokenizer_source or f"{model_path}::tokenizer", f"{model_path}::t5_tokenizer" + return qwen_tokenizer_source or _QWEN_TOKENIZER_SOURCE, _T5_TOKENIZER_SOURCE def load_text_tokenizer(self): qwen_source, t5_source = self._prompt_tokenizer_sources() @@ -313,8 +314,10 @@ def load_text_tokenizer(self): def load_text_encoder(self, move_to_device: bool = True): if self.text_encoders is not None and len(self.text_encoders) > 0: return + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") model_path = ( - getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) + qwen_text_encoder_path + or getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) or self.config.pretrained_model_name_or_path ) revision = getattr(self.config, "text_encoder_revision", None) or getattr(self.config, "revision", None) @@ -324,6 +327,34 @@ def load_text_encoder(self, move_to_device: bool = True): execution_device=self.accelerator.device.type, ) load_device = self.accelerator.device.type if move_to_device else "cpu" + if qwen_text_encoder_path: + if os.path.isfile(model_path): + text_encoder = load_text_encoder_single_file( + file_path=model_path, + device=load_device, + dtype=dtype, + ) + else: + load_kwargs = { + "pretrained_model_name_or_path": model_path, + "revision": revision, + "torch_dtype": dtype, + "local_files_only": bool(getattr(self.config, "local_files_only", False)), + "cache_dir": getattr(self.config, "cache_dir", None), + "force_download": bool(getattr(self.config, "force_download", False)), + } + self._add_hf_token_kwarg(load_kwargs) + text_encoder = Qwen3Model.from_pretrained(**load_kwargs) + text_encoder.eval().requires_grad_(False) + text_encoder.to(device=load_device, dtype=dtype) + self.text_encoders = [text_encoder] + self.text_encoder = text_encoder + self.text_encoder_1 = text_encoder + if not move_to_device: + text_encoder.to("cpu") + if getattr(self, "prompt_tokenizer", None) is None: + self.load_text_tokenizer() + return if self._uses_diffusers_repo_layout(model_path, component_subfolder="text_encoder"): load_kwargs = { "pretrained_model_name_or_path": model_path, diff --git a/simpletuner/helpers/models/auraflow/controlnet.py b/simpletuner/helpers/models/auraflow/controlnet.py index 49b8cdaaa..feb075744 100644 --- a/simpletuner/helpers/models/auraflow/controlnet.py +++ b/simpletuner/helpers/models/auraflow/controlnet.py @@ -372,13 +372,7 @@ def forward( for block in self.joint_transformer_blocks: if self.training and self.gradient_checkpointing: - def create_custom_forward(module): - def custom_forward(*inputs): - return module(*inputs) - - return custom_forward - - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -386,8 +380,17 @@ def custom_forward(*inputs): checkpoint_fn = torch.utils.checkpoint.checkpoint ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + def custom_forward(hidden_states, encoder_hidden_states, temb, checkpoint_block=block): + return checkpoint_block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + attention_kwargs=attention_kwargs, + ) + encoder_hidden_states, hidden_states = checkpoint_fn( - create_custom_forward(block), + custom_forward, hidden_states, encoder_hidden_states, temb, @@ -410,13 +413,7 @@ def custom_forward(*inputs): for block in self.single_transformer_blocks: if self.training and self.gradient_checkpointing: - def create_custom_forward(module): - def custom_forward(*inputs): - return module(*inputs) - - return custom_forward - - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -424,8 +421,16 @@ def custom_forward(*inputs): checkpoint_fn = torch.utils.checkpoint.checkpoint ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + def custom_forward(hidden_states, temb, checkpoint_block=block): + return checkpoint_block( + hidden_states=hidden_states, + temb=temb, + attention_kwargs=attention_kwargs, + ) + combined_hidden_states = checkpoint_fn( - create_custom_forward(block), + custom_forward, combined_hidden_states, temb, **ckpt_kwargs, diff --git a/simpletuner/helpers/models/auraflow/transformer.py b/simpletuner/helpers/models/auraflow/transformer.py index 0ff743712..704205b3f 100644 --- a/simpletuner/helpers/models/auraflow/transformer.py +++ b/simpletuner/helpers/models/auraflow/transformer.py @@ -23,6 +23,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.packed_attention_processors import PackedAuraFlowAttnProcessor2_0 from simpletuner.helpers.training.tread import TREADRouter @@ -493,6 +494,7 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" total_layers = num_mmdit_layers + num_single_dit_layers @@ -512,6 +514,9 @@ def set_gradient_checkpointing_interval(self, interval: int): """ self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend @@ -824,17 +829,16 @@ def _to_pos(idx): if ( self.training - and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) + and torch.is_grad_enabled() + and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) ): - def create_custom_forward(module): - def custom_forward(*inputs): - return module(*inputs) - - return custom_forward - - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -842,8 +846,18 @@ def custom_forward(*inputs): checkpoint_fn = torch.utils.checkpoint.checkpoint ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + def custom_forward(hidden_states, encoder_hidden_states, temb, context_temb, checkpoint_block=block): + return checkpoint_block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + context_temb=context_temb, + attention_kwargs=attention_kwargs, + ) + encoder_hidden_states, hidden_states = checkpoint_fn( - create_custom_forward(block), + custom_forward, hidden_states, encoder_hidden_states, temb, @@ -947,20 +961,16 @@ def custom_forward(*inputs): if ( self.training - and self.gradient_checkpointing - and ( - self.gradient_checkpointing_interval is None - or index_block % self.gradient_checkpointing_interval == 0 + and torch.is_grad_enabled() + and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, ) ): - def create_custom_forward(module): - def custom_forward(*inputs): - return module(*inputs) - - return custom_forward - - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -968,8 +978,16 @@ def custom_forward(*inputs): checkpoint_fn = torch.utils.checkpoint.checkpoint ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + def custom_forward(hidden_states, temb, checkpoint_block=block): + return checkpoint_block( + hidden_states=hidden_states, + temb=temb, + attention_kwargs=attention_kwargs, + ) + combined_hidden_states = checkpoint_fn( - create_custom_forward(block), + custom_forward, combined_hidden_states, single_temb, **ckpt_kwargs, diff --git a/simpletuner/helpers/models/boogu_image/model.py b/simpletuner/helpers/models/boogu_image/model.py index 4b789e0c6..473489878 100644 --- a/simpletuner/helpers/models/boogu_image/model.py +++ b/simpletuner/helpers/models/boogu_image/model.py @@ -16,8 +16,6 @@ get_sdnq_presets, get_torchao_presets, ) -from simpletuner.helpers.models.common import ImageModelFoundation, ModelTypes, PipelineTypes, PredictionTypes -from simpletuner.helpers.models.flux.model import Flux from simpletuner.helpers.models.boogu_image.pipeline import BooguImagePipeline from simpletuner.helpers.models.boogu_image.pipeline_edit import BooguImageEditPipeline from simpletuner.helpers.models.boogu_image.pipeline_img2img import BooguImageImg2ImgPipeline @@ -26,6 +24,8 @@ FlowMatchEulerDiscreteScheduler, ) from simpletuner.helpers.models.boogu_image.transformer import BooguImageTransformer2DModel +from simpletuner.helpers.models.common import ImageModelFoundation, ModelTypes, PipelineTypes, PredictionTypes +from simpletuner.helpers.models.flux.model import Flux from simpletuner.helpers.models.registry import ModelRegistry from simpletuner.helpers.training.deepspeed import deepspeed_zero_init_disabled_context_manager @@ -188,8 +188,8 @@ def text_embed_cache_key(self): def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_kwargs = { "pretrained_model_name_or_path": processor_path, "subfolder": processor_subfolder, diff --git a/simpletuner/helpers/models/boogu_image/transformer.py b/simpletuner/helpers/models/boogu_image/transformer.py index 79e28d552..154eb6328 100644 --- a/simpletuner/helpers/models/boogu_image/transformer.py +++ b/simpletuner/helpers/models/boogu_image/transformer.py @@ -28,6 +28,8 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers from einops import rearrange +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block + from .attention_processor import BooguImageAttnProcessor, BooguImageDoubleStreamSelfAttnProcessor from .block_lumina2 import ( Lumina2CombinedTimestepCaptionEmbedding, @@ -842,6 +844,8 @@ def __init__( self.image_index_embedding = nn.Parameter(torch.randn(5, hidden_size)) # support max 5 ref images self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.initialize_weights() @@ -858,6 +862,12 @@ def __init__( self.layers = list(self.double_stream_layers) + list(self.single_stream_layers) + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def initialize_weights(self) -> None: """ Initialize the weights of the model. @@ -897,6 +907,10 @@ def img_patch_embed_and_refine( hidden_states = self.x_embedder(hidden_states) ref_image_hidden_states = self.ref_image_patch_embedder(ref_image_hidden_states) + if hidden_states.requires_grad: + hidden_states = hidden_states.clone() + if ref_image_hidden_states.requires_grad: + ref_image_hidden_states = ref_image_hidden_states.clone() for i in range(batch_size): shift = 0 @@ -1272,7 +1286,12 @@ def forward( else: layer.enable_taylorseer = False - if torch.is_grad_enabled() and self.gradient_checkpointing: + if torch.is_grad_enabled() and should_checkpoint_block( + layer_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): img_hidden_states, instruct_hidden_states = self._gradient_checkpointing_func( layer, img_hidden_states, @@ -1355,7 +1374,12 @@ def forward( layer.enable_taylorseer = True self.current["layer"] = self.num_double_stream_layers + layer_idx - if torch.is_grad_enabled() and self.gradient_checkpointing: + if torch.is_grad_enabled() and should_checkpoint_block( + self.num_double_stream_layers + layer_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): hidden_states = self._gradient_checkpointing_func( layer, hidden_states, joint_attention_mask, rotary_emb, temb ) diff --git a/simpletuner/helpers/models/chroma/transformer.py b/simpletuner/helpers/models/chroma/transformer.py index 283a1f839..9ede58903 100644 --- a/simpletuner/helpers/models/chroma/transformer.py +++ b/simpletuner/helpers/models/chroma/transformer.py @@ -31,7 +31,9 @@ from simpletuner.helpers.models.flux.attention import FluxFusedFlashAttnProcessor3 from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.attention_backend import AttentionBackendController +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter @@ -350,6 +352,24 @@ def __init__( pre_only=True, ) + def _ffn_forward( + self, + residual: torch.Tensor, + norm_hidden_states: torch.Tensor, + attn_output: torch.Tensor, + gate: torch.Tensor, + ) -> torch.Tensor: + mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) + hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) + if gate.ndim == 2: + gate = gate.unsqueeze(1) + hidden_states = gate * self.proj_out(hidden_states) + hidden_states = residual + hidden_states + if hidden_states.dtype == torch.float16: + hidden_states = hidden_states.clip(-65504, 65504) + + return hidden_states + def forward( self, hidden_states: torch.Tensor, @@ -357,10 +377,12 @@ def forward( image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, attention_mask: Optional[torch.Tensor] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> torch.Tensor: residual = hidden_states norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)) joint_attention_kwargs = joint_attention_kwargs or {} if attention_mask is not None: @@ -394,12 +416,13 @@ def forward( except Exception: logger.debug("ChromaFluxSingleTransformerBlock failed to publish QK-Clip logits.", exc_info=True) - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - **joint_attention_kwargs, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output = self.attn( + hidden_states=norm_hidden_states, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + **joint_attention_kwargs, + ) publish_attention_max_logits( getattr(self.attn, "last_query", None) if hasattr(self.attn, "last_query") else None, getattr(self.attn, "last_key", None) if hasattr(self.attn, "last_key") else None, @@ -408,15 +431,17 @@ def forward( getattr(self.attn, "to_k", None) and self.attn.to_k.weight, ) - hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) - if gate.ndim == 2: - gate = gate.unsqueeze(1) - hidden_states = gate * self.proj_out(hidden_states) - hidden_states = residual + hidden_states - if hidden_states.dtype == torch.float16: - hidden_states = hidden_states.clip(-65504, 65504) - - return hidden_states + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + return checkpoint_fn( + self._ffn_forward, + residual, + norm_hidden_states, + attn_output, + gate, + ) + return self._ffn_forward(residual, norm_hidden_states, attn_output, gate) @maybe_allow_in_graph @@ -451,6 +476,43 @@ def __init__( self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) self.ff_context = FeedForward(dim=dim, dim_out=dim, activation_fn="gelu-approximate") + def _ffn_forward( + self, + hidden_states: torch.Tensor, + encoder_hidden_states: torch.Tensor, + shift_mlp: torch.Tensor, + scale_mlp: torch.Tensor, + gate_mlp: torch.Tensor, + c_shift_mlp: torch.Tensor, + c_scale_mlp: torch.Tensor, + c_gate_mlp: torch.Tensor, + ip_attn_output: Optional[torch.Tensor] = None, + ) -> Tuple[torch.Tensor, torch.Tensor]: + norm_hidden_states = self.norm2(hidden_states) + if scale_mlp.ndim == 2: + norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + else: + norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp + + ff_output = self.ff(norm_hidden_states) + if gate_mlp.ndim == 2: + gate_mlp = gate_mlp.unsqueeze(1) + hidden_states = hidden_states + gate_mlp * ff_output + if ip_attn_output is not None: + hidden_states = hidden_states + ip_attn_output + + norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) + norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] + + context_ff_output = self.ff_context(norm_encoder_hidden_states) + if c_gate_mlp.ndim == 2: + c_gate_mlp = c_gate_mlp.unsqueeze(1) + encoder_hidden_states = encoder_hidden_states + c_gate_mlp * context_ff_output + if encoder_hidden_states.dtype == torch.float16: + encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) + + return encoder_hidden_states, hidden_states + def forward( self, hidden_states: torch.Tensor, @@ -460,6 +522,9 @@ def forward( image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, attention_mask: Optional[torch.Tensor] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=image_temb) @@ -497,14 +562,16 @@ def forward( except Exception: logger.debug("ChromaTransformerBlock failed to publish QK-Clip logits.", exc_info=True) - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - **joint_attention_kwargs, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attention_outputs = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + **joint_attention_kwargs, + ) + ip_attn_output = None if len(attention_outputs) == 2: attn_output, context_attn_output = attention_outputs elif len(attention_outputs) == 3: @@ -515,37 +582,38 @@ def forward( attn_output = gate_msa * attn_output hidden_states = hidden_states + attn_output - norm_hidden_states = self.norm2(hidden_states) - if scale_mlp.ndim == 2: - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - else: - norm_hidden_states = norm_hidden_states * (1 + scale_mlp) + shift_mlp - - ff_output = self.ff(norm_hidden_states) - if gate_mlp.ndim == 2: - gate_mlp = gate_mlp.unsqueeze(1) - ff_output = gate_mlp * ff_output - - hidden_states = hidden_states + ff_output - if len(attention_outputs) == 3: - hidden_states = hidden_states + ip_attn_output - if c_gate_msa.ndim == 2: c_gate_msa = c_gate_msa.unsqueeze(1) context_attn_output = c_gate_msa * context_attn_output encoder_hidden_states = encoder_hidden_states + context_attn_output - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None] - - context_ff_output = self.ff_context(norm_encoder_hidden_states) - if c_gate_mlp.ndim == 2: - c_gate_mlp = c_gate_mlp.unsqueeze(1) - encoder_hidden_states = encoder_hidden_states + c_gate_mlp * context_ff_output - if encoder_hidden_states.dtype == torch.float16: - encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504) - - return encoder_hidden_states, hidden_states + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + return checkpoint_fn( + self._ffn_forward, + hidden_states, + encoder_hidden_states, + shift_mlp, + scale_mlp, + gate_mlp, + c_shift_mlp, + c_scale_mlp, + c_gate_mlp, + ip_attn_output, + use_reentrant=False, + ) + return self._ffn_forward( + hidden_states, + encoder_hidden_states, + shift_mlp, + scale_mlp, + gate_mlp, + c_shift_mlp, + c_scale_mlp, + c_gate_mlp, + ip_attn_output, + ) class ChromaTransformer2DModel( @@ -562,6 +630,8 @@ class ChromaTransformer2DModel( """ _supports_gradient_checkpointing = True + _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _no_split_modules = ["ChromaTransformerBlock", "ChromaSingleTransformerBlock"] _repeated_blocks = ["ChromaTransformerBlock", "ChromaSingleTransformerBlock"] _skip_layerwise_casting_patterns = ["pos_embed", "norm"] @@ -679,6 +749,8 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None + self.gradient_checkpointing_offload_attention = False self.gradient_checkpointing_backend = "torch" total_layers = num_layers + num_single_layers @@ -708,6 +780,12 @@ def set_gradient_checkpointing_interval(self, interval: int): """ self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend @@ -949,26 +1027,66 @@ def _to_pos(idx): use_checkpoint = ( torch.is_grad_enabled() and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) + and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) ) if use_checkpoint: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = self._gradient_checkpointing_func - encoder_hidden_states, hidden_states = checkpoint_fn( - block, - hidden_states, - encoder_hidden_states, - image_temb, - text_temb, - current_rope, - attention_mask, - ) + if self.gradient_checkpointing_backend.endswith("-ffn"): + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + image_temb=image_temb, + text_temb=text_temb, + image_rotary_emb=current_rope, + attention_mask=attention_mask, + joint_attention_kwargs=joint_attention_kwargs, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_checkpointed_block( + checkpoint_hidden_states, + checkpoint_encoder_hidden_states, + checkpoint_image_temb, + checkpoint_text_temb, + checkpoint_rope, + checkpoint_attention_mask, + checkpoint_block=block, + ): + return checkpoint_block( + hidden_states=checkpoint_hidden_states, + encoder_hidden_states=checkpoint_encoder_hidden_states, + image_temb=checkpoint_image_temb, + text_temb=checkpoint_text_temb, + image_rotary_emb=checkpoint_rope, + attention_mask=checkpoint_attention_mask, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + encoder_hidden_states, hidden_states = checkpoint_fn( + run_checkpointed_block, + hidden_states, + encoder_hidden_states, + image_temb, + text_temb, + current_rope, + attention_mask, + ) else: encoder_hidden_states, hidden_states = block( @@ -979,6 +1097,7 @@ def _to_pos(idx): image_rotary_emb=current_rope, attention_mask=attention_mask, joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(block, "fuser"): @@ -1110,23 +1229,56 @@ def _to_pos(idx): use_checkpoint = ( torch.is_grad_enabled() and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) + and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) ) if use_checkpoint: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = self._gradient_checkpointing_func - hidden_states = checkpoint_fn( - block, - hidden_states, - temb, - current_rope, - ) + if self.gradient_checkpointing_backend.endswith("-ffn"): + hidden_states = block( + hidden_states=hidden_states, + temb=temb, + image_rotary_emb=current_rope, + attention_mask=attention_mask, + joint_attention_kwargs=joint_attention_kwargs, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_checkpointed_block( + checkpoint_hidden_states, + checkpoint_temb, + checkpoint_rope, + checkpoint_block=block, + ): + return checkpoint_block( + hidden_states=checkpoint_hidden_states, + temb=checkpoint_temb, + image_rotary_emb=checkpoint_rope, + attention_mask=attention_mask, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + hidden_states = checkpoint_fn( + run_checkpointed_block, + hidden_states, + temb, + current_rope, + ) else: hidden_states = block( @@ -1135,6 +1287,7 @@ def _to_pos(idx): image_rotary_emb=current_rope, attention_mask=attention_mask, joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(block, "fuser"): diff --git a/simpletuner/helpers/models/common.py b/simpletuner/helpers/models/common.py index 5253afd7d..f887cd41e 100644 --- a/simpletuner/helpers/models/common.py +++ b/simpletuner/helpers/models/common.py @@ -2517,6 +2517,84 @@ def _is_gemma_component(component_cls) -> bool: component_name = getattr(component_cls, "__name__", "") return "gemma" in component_name.lower() + @staticmethod + def _component_value_mentions_qwen(value) -> bool: + if value is None: + return False + if isinstance(value, str): + return "qwen" in value.lower() + component_name = getattr(value, "__name__", "") + return "qwen" in component_name.lower() + + @staticmethod + def _optional_model_path(value) -> Optional[str]: + if value is None or value is False: + return None + if isinstance(value, os.PathLike): + return os.fspath(value) + if isinstance(value, str): + value = value.strip() + return value or None + return None + + def _get_optional_config_model_path(self, name: str) -> Optional[str]: + return self._optional_model_path(getattr(self.config, name, None)) + + def _is_qwen_text_encoder_config(self, text_encoder_config: dict) -> bool: + return any( + self._component_value_mentions_qwen(text_encoder_config.get(key)) + for key in ("name", "model", "tokenizer", "path") + ) + + def _qwen_text_encoder_config_count(self) -> int: + text_encoder_configuration = getattr(self, "TEXT_ENCODER_CONFIGURATION", None) or {} + return sum( + 1 + for text_encoder_config in text_encoder_configuration.values() + if self._is_qwen_text_encoder_config(text_encoder_config) + ) + + def _warn_qwen_text_encoder_override_ignored(self, qwen_count: int) -> None: + if getattr(self, "_qwen_text_encoder_override_warning_emitted", False): + return + logger.warning( + "Ignoring qwen_text_encoder_model_name_or_path for %s because this model defines %s Qwen text encoders.", + self.NAME, + qwen_count, + ) + self._qwen_text_encoder_override_warning_emitted = True + + def _uses_qwen_text_encoder_override(self, text_encoder_config: dict) -> bool: + qwen_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + if not qwen_path or not self._is_qwen_text_encoder_config(text_encoder_config): + return False + qwen_count = self._qwen_text_encoder_config_count() + if qwen_count == 1: + return True + if qwen_count > 1: + self._warn_qwen_text_encoder_override_ignored(qwen_count) + return False + + def _resolve_text_encoder_subfolder(self, text_encoder_config: dict, key: str, default=None): + if self._uses_qwen_text_encoder_override(text_encoder_config): + return None + return text_encoder_config.get(key, default) + + def _resolve_qwen_processor_path(self, default_path: str) -> str: + return ( + self._get_optional_config_model_path("processor_pretrained_model_name_or_path") + or self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + or default_path + ) + + def _resolve_qwen_processor_subfolder(self, default_subfolder: Optional[str]) -> Optional[str]: + processor_subfolder = getattr(self.config, "processor_subfolder", None) + if isinstance(processor_subfolder, str): + return processor_subfolder + if self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path"): + return None + return default_subfolder + def _resolve_text_encoder_path(self, text_encoder_config: dict) -> str: text_encoder_path = get_model_config_path(self.config.model_family, self.config.pretrained_model_name_or_path) config_path = text_encoder_config.get("path", None) @@ -2525,6 +2603,9 @@ def _resolve_text_encoder_path(self, text_encoder_config: dict) -> str: gemma_path = getattr(self.config, "pretrained_gemma_model_name_or_path", None) if gemma_path and self._is_gemma_component(text_encoder_config.get("model")): text_encoder_path = gemma_path + qwen_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + if qwen_path and self._uses_qwen_text_encoder_override(text_encoder_config): + text_encoder_path = qwen_path return text_encoder_path def load_text_tokenizer(self): @@ -2551,7 +2632,11 @@ def load_text_tokenizer(self): for attr_name, text_encoder_config in self.TEXT_ENCODER_CONFIGURATION.items(): tokenizer_idx += 1 tokenizer_cls = text_encoder_config.get("tokenizer") - tokenizer_kwargs["subfolder"] = text_encoder_config.get("tokenizer_subfolder", "tokenizer") + tokenizer_kwargs["subfolder"] = self._resolve_text_encoder_subfolder( + text_encoder_config, + "tokenizer_subfolder", + "tokenizer", + ) tokenizer_kwargs["use_fast"] = text_encoder_config.get("use_fast", False) tokenizer_kwargs["pretrained_model_name_or_path"] = self._resolve_text_encoder_path(text_encoder_config) logger.info(f"Loading tokenizer {tokenizer_idx}: {tokenizer_cls.__name__} with args: {tokenizer_kwargs}") @@ -2673,7 +2758,12 @@ def load_text_encoder(self, move_to_device: bool = True): "pretrained_model_name_or_path": text_encoder_path, "variant": self.config.variant, "revision": self.config.revision, - "subfolder": text_encoder_config.get("subfolder", "text_encoder") or "", + "subfolder": self._resolve_text_encoder_subfolder( + text_encoder_config, + "subfolder", + "text_encoder", + ) + or "", **extra_kwargs, } accelerator = getattr(self, "accelerator", None) @@ -3176,16 +3266,61 @@ def _load_model(load_kwargs: dict): if checkpoint_backend in {"unsloth", "unsloth-ffn"}: logger.info("Using Unsloth-style gradient checkpointing (CPU offload)") + offload_attention = bool(getattr(self.config, "gradient_checkpointing_offload_attention", False)) + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + normalize_activation_offload_pin_memory_max_buckets, + set_activation_offload_pin_memory_max_buckets, + set_activation_offload_prefetch_enabled, + ) + + offload_pin_memory_max_buckets = normalize_activation_offload_pin_memory_max_buckets( + getattr(self.config, "gradient_checkpointing_offload_pin_memory_max_buckets", 12) + ) + set_activation_offload_pin_memory_max_buckets(offload_pin_memory_max_buckets) + set_activation_offload_prefetch_enabled(bool(getattr(self.config, "gradient_checkpointing_offload_prefetch", False))) + if offload_attention: + trained_component = self.unwrap_model(model=self.model) if self.model is not None else None + if trained_component is None or not getattr(trained_component, "_supports_attention_activation_offload", False): + raise ValueError( + "--gradient_checkpointing_offload_attention requires a model with an attention/FFN checkpointing " + f"boundary, but {self.config.model_family} does not expose it." + ) + logger.info("Using attention activation offload for gradient checkpointing") + if self.config.gradient_checkpointing_interval is not None and self.config.gradient_checkpointing_interval > 1: if self.model is not None and hasattr(self.model, "set_gradient_checkpointing_interval"): logger.info("Setting gradient checkpointing interval..") self.unwrap_model(model=self.model).set_gradient_checkpointing_interval( int(self.config.gradient_checkpointing_interval) ) + else: + logger.warning( + "--gradient_checkpointing_interval=%s was requested, but %s does not expose interval " + "checkpointing support; the value will be ignored.", + self.config.gradient_checkpointing_interval, + self.config.model_family, + ) + + gradient_checkpointing_segment_stride = getattr(self.config, "gradient_checkpointing_segment_stride", None) + if gradient_checkpointing_segment_stride is not None: + if self.model is not None and hasattr(self.model, "set_gradient_checkpointing_segment_stride"): + logger.info("Setting gradient checkpointing segment stride..") + self.unwrap_model(model=self.model).set_gradient_checkpointing_segment_stride( + int(gradient_checkpointing_segment_stride) + ) + else: + logger.warning( + "--gradient_checkpointing_segment_stride=%s was requested, but %s does not expose segmented " + "checkpoint stride support; the value will be ignored.", + gradient_checkpointing_segment_stride, + self.config.model_family, + ) # Set gradient checkpointing backend on model if supported if self.model is not None and hasattr(self.model, "set_gradient_checkpointing_backend"): self.unwrap_model(model=self.model).set_gradient_checkpointing_backend(checkpoint_backend) + if self.model is not None and hasattr(self.model, "set_gradient_checkpointing_offload_attention"): + self.unwrap_model(model=self.model).set_gradient_checkpointing_offload_attention(offload_attention) self.fuse_qkv_projections() self.post_model_load_setup() diff --git a/simpletuner/helpers/models/cosmos/transformer.py b/simpletuner/helpers/models/cosmos/transformer.py index 927ec1b53..11186cbae 100644 --- a/simpletuner/helpers/models/cosmos/transformer.py +++ b/simpletuner/helpers/models/cosmos/transformer.py @@ -41,6 +41,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import MutableModuleList, PatchableModule @@ -677,6 +678,8 @@ def __init__( ) self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, blocks_to_swap=musubi_blocks_to_swap, @@ -697,6 +700,12 @@ def set_router(self, router: TREADRouter, routes: Optional[List[Dict]] = None): self._tread_router = router self._tread_routes = routes + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def forward( self, hidden_states: torch.Tensor, @@ -796,7 +805,6 @@ def forward( # 3. Patchify input p_t, p_h, p_w = self.config.patch_size - expected_num_frames = num_frames // p_t expected_height = height // p_h expected_width = width // p_w hidden_states = self.patch_embed(hidden_states) @@ -889,7 +897,16 @@ def forward( break if musubi_offload_active and musubi_manager.is_managed_block(bid): musubi_manager.stream_in(block, hidden_states.device) - if torch.is_grad_enabled() and self.gradient_checkpointing: + if ( + grad_enabled + and self.gradient_checkpointing + and should_checkpoint_block( + bid, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + ): hidden_states = self._gradient_checkpointing_func( block, hidden_states, diff --git a/simpletuner/helpers/models/cosmos3/transformer.py b/simpletuner/helpers/models/cosmos3/transformer.py index e2ecc7536..ebd82b720 100644 --- a/simpletuner/helpers/models/cosmos3/transformer.py +++ b/simpletuner/helpers/models/cosmos3/transformer.py @@ -27,6 +27,7 @@ from diffusers.utils import BaseOutput, logging from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block logger = logging.get_logger(__name__) @@ -666,6 +667,7 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_hidden_layers, blocks_to_swap=musubi_blocks_to_swap, @@ -689,10 +691,18 @@ def unfuse_qkv_projections(self) -> None: def set_gradient_checkpointing_interval(self, interval: int): self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def _should_gradient_checkpoint_layer(self, layer_idx: int) -> bool: if not torch.is_grad_enabled() or not self.gradient_checkpointing: return False - return self.gradient_checkpointing_interval is None or layer_idx % self.gradient_checkpointing_interval == 0 + return should_checkpoint_block( + layer_idx, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) @staticmethod def _stream_in_for_checkpoint_recompute(musubi_manager, layer_idx: int, layer: nn.Module, device: torch.device) -> None: @@ -772,12 +782,9 @@ def _unpatchify_and_unpack_latents( h_patches = h_padded // p w_patches = w_padded // p t_n = len(noisy_frame_indexes) - output_tensor = torch.zeros( - (latent_channel, t_c, h_orig, w_orig), - device=packed_mse_preds.device, - dtype=packed_mse_preds.dtype, - ) num_patches = t_n * h_patches * w_patches + output_dtype = packed_mse_preds.dtype + latent = None if num_patches > 0: end_idx = start_idx + num_patches latent_patches = packed_mse_preds[start_idx:end_idx] @@ -785,6 +792,14 @@ def _unpatchify_and_unpack_latents( latent = torch.einsum("thwpqc->cthpwq", latent_patches) latent = latent.reshape(latent_channel, t_n, h_patches * p, w_patches * p) latent = latent[:, :, :h_orig, :w_orig] + output_dtype = latent.dtype + + output_tensor = torch.zeros( + (latent_channel, t_c, h_orig, w_orig), + device=packed_mse_preds.device, + dtype=output_dtype, + ) + if latent is not None: output_tensor[:, noisy_frame_indexes] = latent start_idx = end_idx unpatchified_latents.append(output_tensor.unsqueeze(0)) diff --git a/simpletuner/helpers/models/ernie/transformer.py b/simpletuner/helpers/models/ernie/transformer.py index 91b2471b6..e41542ca2 100644 --- a/simpletuner/helpers/models/ernie/transformer.py +++ b/simpletuner/helpers/models/ernie/transformer.py @@ -15,6 +15,7 @@ shard_cp_tensor, unshard_cp_tensor, ) +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.tread import TREADRouter logger = logging.getLogger(__name__) @@ -66,6 +67,9 @@ def set_router(self, router: TREADRouter, routes: Optional[List[Dict[str, Any]]] def set_gradient_checkpointing_interval(self, interval: int): self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend @@ -217,9 +221,14 @@ def apply_layer(layer_idx, layer, batch_first_hidden_states, rotary_emb, attn_ma if ( torch.is_grad_enabled() and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or layer_idx % self.gradient_checkpointing_interval == 0) + and should_checkpoint_block( + layer_idx, + True, + self.gradient_checkpointing_interval, + getattr(self, "gradient_checkpointing_segment_stride", None), + ) ): - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint sequence_first_hidden_states = offloaded_checkpoint( diff --git a/simpletuner/helpers/models/ernie/transformer_diffusers.py b/simpletuner/helpers/models/ernie/transformer_diffusers.py index 52f6ca209..15dfb5c6b 100644 --- a/simpletuner/helpers/models/ernie/transformer_diffusers.py +++ b/simpletuner/helpers/models/ernie/transformer_diffusers.py @@ -502,7 +502,7 @@ def forward( and self.gradient_checkpointing and (self.gradient_checkpointing_interval is None or layer_idx % self.gradient_checkpointing_interval == 0) ): - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint x = offloaded_checkpoint(layer, x, rotary_pos_emb, temb, attention_mask, use_reentrant=False) diff --git a/simpletuner/helpers/models/flux/transformer.py b/simpletuner/helpers/models/flux/transformer.py index b7d9270de..66d2833ee 100644 --- a/simpletuner/helpers/models/flux/transformer.py +++ b/simpletuner/helpers/models/flux/transformer.py @@ -46,7 +46,9 @@ from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.attention_backend import maybe_metal_flash_rope_attention from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import CallableDict, MutableModuleList, PatchableModule @@ -476,6 +478,7 @@ def forward( attention_mask: Optional[torch.Tensor] = None, checkpoint_ffn: bool = False, checkpoint_fn: Any | None = None, + offload_attention: bool = False, ): residual = hidden_states norm_hidden_states, gate = _flux_apply_ada_layer_norm_zero_single(self.norm, hidden_states, temb) @@ -486,11 +489,12 @@ def forward( attention_mask, ) - attn_output = self.attn( - hidden_states=norm_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output = self.attn( + hidden_states=norm_hidden_states, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + ) if checkpoint_ffn: if checkpoint_fn is None: @@ -610,6 +614,7 @@ def forward( attention_mask: Optional[torch.Tensor] = None, checkpoint_ffn: bool = False, checkpoint_fn: Any | None = None, + offload_attention: bool = False, ): if context_temb is None: context_temb = temb @@ -628,12 +633,13 @@ def forward( ) # Attention. - attention_outputs = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - attention_mask=attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attention_outputs = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + ) ip_attn_output = None if len(attention_outputs) == 2: attn_output, context_attn_output = attention_outputs @@ -701,6 +707,7 @@ class FluxTransformer2DModel(PatchableModule, ModelMixin, ConfigMixin, PeftAdapt _supports_gradient_checkpointing = True _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True # Hint FSDP auto wrap policy to shard per transformer block. _no_split_modules = ["FluxTransformerBlock", "FluxSingleTransformerBlock"] _cp_plan = { @@ -807,7 +814,9 @@ def __init__( self.gradient_checkpointing = False # optional interval for gradient checkpointing self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False total_layers = num_layers + num_single_layers self._musubi_block_swap = MusubiBlockSwapManager.build( depth=total_layers, @@ -819,9 +828,15 @@ def __init__( def set_gradient_checkpointing_interval(self, value: int): self.gradient_checkpointing_interval = value + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): self._tread_router = router self._tread_routes = routes @@ -1124,8 +1139,75 @@ def _to_pos(idx): if hasattr(self, "position_net") and grounding_kwargs is not None: grounding_objs = self.position_net(**grounding_kwargs) + segment_size = self.gradient_checkpointing_interval + has_te_checkpoint_context = any( + getattr(block, "_simpletuner_te_checkpoint_context_fn", None) is not None for block in combined_blocks + ) + use_segmented_checkpointing = ( + self.training + and self.gradient_checkpointing + and segment_size is not None + and segment_size > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not musubi_offload_active + and not has_te_checkpoint_context + and grounding_objs is None + and hidden_states_buffer is None + ) + segmented_checkpoint_fn = None + segmented_checkpoint_kwargs: Dict[str, Any] = {} + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + segmented_checkpoint_kwargs = {"use_reentrant": False} + capture_idx = 0 for index_block, block in enumerate(self.transformer_blocks): + run_gap_eagerly = False + if use_segmented_checkpointing and controlnet_block_samples is None: + segment_stride = self.gradient_checkpointing_segment_stride or segment_size + segment_offset = index_block % segment_stride + if segment_offset < segment_size and segment_offset != 0: + continue + if segment_offset >= segment_size: + run_gap_eagerly = True + else: + segment_blocks = list(self.transformer_blocks[index_block : index_block + segment_size]) + + def run_double_block( + _relative_index, + segment_block, + segment_hidden_states, + segment_encoder_hidden_states, + ): + next_encoder_hidden_states, next_hidden_states = segment_block( + hidden_states=segment_hidden_states, + encoder_hidden_states=segment_encoder_hidden_states, + temb=temb_img, + context_temb=temb_txt, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + return next_hidden_states, next_encoder_hidden_states + + hidden_states, encoder_hidden_states = checkpoint_sequential_state( + segment_blocks, + len(segment_blocks), + (hidden_states, encoder_hidden_states), + run_double_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + ) + global_idx += len(segment_blocks) + capture_idx += len(segment_blocks) + continue + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): musubi_manager.stream_in(block, hidden_states.device) # TREAD: START a route? @@ -1158,16 +1240,14 @@ def _to_pos(idx): if ( self.training and self.gradient_checkpointing + and not run_gap_eagerly and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) ): checkpoint_ffn = self.gradient_checkpointing_backend.endswith("-ffn") - def create_custom_forward(module, return_dict=None): + def create_custom_forward(module): def custom_forward(*inputs): - if return_dict is not None: - return module(*inputs, return_dict=return_dict) - else: - return module(*inputs) + return module(*inputs, offload_attention=self.gradient_checkpointing_offload_attention) return custom_forward @@ -1190,6 +1270,7 @@ def custom_forward(*inputs): attention_mask=attention_mask, checkpoint_ffn=True, checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, ) else: if is_torch_version(">=", "1.11.0"): @@ -1213,6 +1294,7 @@ def custom_forward(*inputs): context_temb=temb_txt, image_rotary_emb=current_rope, attention_mask=attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(block, "fuser"): @@ -1251,6 +1333,38 @@ def custom_forward(*inputs): txt_len = encoder_hidden_states.shape[1] for index_block, block in enumerate(self.single_transformer_blocks): + run_gap_eagerly = False + if use_segmented_checkpointing and controlnet_single_block_samples is None: + segment_stride = self.gradient_checkpointing_segment_stride or segment_size + segment_offset = index_block % segment_stride + if segment_offset < segment_size and segment_offset != 0: + continue + if segment_offset >= segment_size: + run_gap_eagerly = True + else: + segment_blocks = list(self.single_transformer_blocks[index_block : index_block + segment_size]) + + def run_single_block(_relative_index, segment_block, segment_hidden_states): + return segment_block( + hidden_states=segment_hidden_states, + temb=temb_single, + image_rotary_emb=image_rotary_emb, + attention_mask=attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + (hidden_states,) = checkpoint_sequential_state( + segment_blocks, + len(segment_blocks), + (hidden_states,), + run_single_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + ) + global_idx += len(segment_blocks) + capture_idx += len(segment_blocks) + continue + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): musubi_manager.stream_in(block, hidden_states.device) # TREAD: START? (operate on *image* tokens only) @@ -1291,19 +1405,14 @@ def custom_forward(*inputs): if ( self.training and self.gradient_checkpointing - or ( - self.gradient_checkpointing_interval is not None - and index_block % self.gradient_checkpointing_interval == 0 - ) + and not run_gap_eagerly + and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) ): checkpoint_ffn = self.gradient_checkpointing_backend.endswith("-ffn") - def create_custom_forward(module, return_dict=None): + def create_custom_forward(module): def custom_forward(*inputs): - if return_dict is not None: - return module(*inputs, return_dict=return_dict) - else: - return module(*inputs) + return module(*inputs, offload_attention=self.gradient_checkpointing_offload_attention) return custom_forward @@ -1324,6 +1433,7 @@ def custom_forward(*inputs): attention_mask=attention_mask, checkpoint_ffn=True, checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, ) else: if is_torch_version(">=", "1.11.0"): @@ -1343,6 +1453,7 @@ def custom_forward(*inputs): temb=temb_single, image_rotary_emb=current_rope, attention_mask=attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(block, "fuser"): diff --git a/simpletuner/helpers/models/flux2/model.py b/simpletuner/helpers/models/flux2/model.py index 0f8d0447a..b470dc690 100644 --- a/simpletuner/helpers/models/flux2/model.py +++ b/simpletuner/helpers/models/flux2/model.py @@ -455,7 +455,9 @@ def _load_text_encoder_qwen3(self, move_to_device: bool = True): # For Klein models, text encoder is bundled in the model repo under "text_encoder" subfolder # and tokenizer is in a separate "tokenizer" subfolder model_path = self.config.pretrained_model_name_or_path - text_encoder_path = getattr(self.config, "pretrained_text_encoder_model_name_or_path", None) + text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") or getattr( + self.config, "pretrained_text_encoder_model_name_or_path", None + ) if text_encoder_path is None: text_encoder_path = model_path text_encoder_subfolder = "text_encoder" diff --git a/simpletuner/helpers/models/flux2/transformer.py b/simpletuner/helpers/models/flux2/transformer.py index 99c899337..26ae31d38 100644 --- a/simpletuner/helpers/models/flux2/transformer.py +++ b/simpletuner/helpers/models/flux2/transformer.py @@ -40,7 +40,10 @@ ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.attention_backend import get_packed_attention_backend, maybe_metal_flash_rope_attention +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint from simpletuner.helpers.training.context_parallel_tensors import context_parallel_config, prepare_cp_attention_mask +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -194,6 +197,7 @@ def __call__( encoder_hidden_states: torch.Tensor = None, attention_mask: Optional[torch.Tensor] = None, image_rotary_emb: Optional[torch.Tensor] = None, + offload_attention: bool = False, ) -> torch.Tensor: query, key, value, encoder_query, encoder_key, encoder_value = _get_qkv_projections( attn, hidden_states, encoder_hidden_states @@ -231,15 +235,16 @@ def __call__( hidden_states = None if image_rotary_emb is not None and self._packed_attention_backend is None and not cp_active: - hidden_states = maybe_metal_flash_rope_attention( - query, - key, - value, - image_rotary_emb, - attn_mask=attention_mask, - backend=self._attention_backend, - layout="bshd", - ) + with activation_offload_context(offload_attention, label=f"{attn.__class__.__qualname__}:attention"): + hidden_states = maybe_metal_flash_rope_attention( + query, + key, + value, + image_rotary_emb, + attn_mask=attention_mask, + backend=self._attention_backend, + layout="bshd", + ) if hidden_states is None: if image_rotary_emb is not None: @@ -254,17 +259,21 @@ def __call__( getattr(attn, "to_k", None) and getattr(attn, "to_k", None).weight, ) - if hidden_states is None and self._packed_attention_backend is not None and not cp_active: - hidden_states = _run_packed_qkv_attention(query, key, value, attention_mask, self._packed_attention_backend) - elif hidden_states is None: - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=parallel_config, - ) + if hidden_states is None: + with activation_offload_context(offload_attention, label=f"{attn.__class__.__qualname__}:attention"): + if self._packed_attention_backend is not None and not cp_active: + hidden_states = _run_packed_qkv_attention( + query, key, value, attention_mask, self._packed_attention_backend + ) + else: + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + backend=self._attention_backend, + parallel_config=parallel_config, + ) hidden_states = hidden_states.flatten(2, 3) hidden_states = hidden_states.to(query.dtype) @@ -373,6 +382,7 @@ def __call__( hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, image_rotary_emb: Optional[torch.Tensor] = None, + offload_attention: bool = False, ) -> torch.Tensor: # Parallel in (QKV + MLP in) projection hidden_states = attn.to_qkv_mlp_proj(hidden_states) @@ -403,15 +413,16 @@ def __call__( hidden_states = None if image_rotary_emb is not None and self._packed_attention_backend is None and not cp_active: - hidden_states = maybe_metal_flash_rope_attention( - query, - key, - value, - image_rotary_emb, - attn_mask=attention_mask, - backend=self._attention_backend, - layout="bshd", - ) + with activation_offload_context(offload_attention, label=f"{attn.__class__.__qualname__}:attention"): + hidden_states = maybe_metal_flash_rope_attention( + query, + key, + value, + image_rotary_emb, + attn_mask=attention_mask, + backend=self._attention_backend, + layout="bshd", + ) if hidden_states is None: if image_rotary_emb is not None: @@ -426,17 +437,21 @@ def __call__( getattr(attn, "to_qkv_mlp_proj", None) and attn.to_qkv_mlp_proj.weight, ) - if hidden_states is None and self._packed_attention_backend is not None and not cp_active: - hidden_states = _run_packed_qkv_attention(query, key, value, attention_mask, self._packed_attention_backend) - elif hidden_states is None: - hidden_states = dispatch_attention_fn( - query, - key, - value, - attn_mask=attention_mask, - backend=self._attention_backend, - parallel_config=parallel_config, - ) + if hidden_states is None: + with activation_offload_context(offload_attention, label=f"{attn.__class__.__qualname__}:attention"): + if self._packed_attention_backend is not None and not cp_active: + hidden_states = _run_packed_qkv_attention( + query, key, value, attention_mask, self._packed_attention_backend + ) + else: + hidden_states = dispatch_attention_fn( + query, + key, + value, + attn_mask=attention_mask, + backend=self._attention_backend, + parallel_config=parallel_config, + ) hidden_states = hidden_states.flatten(2, 3) hidden_states = hidden_states.to(query.dtype) @@ -567,6 +582,7 @@ def forward( joint_attention_kwargs: Optional[Dict[str, Any]] = None, split_hidden_states: bool = False, text_seq_len: Optional[int] = None, + offload_attention: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: # If encoder_hidden_states is None, hidden_states is assumed to have encoder_hidden_states already # concatenated @@ -583,6 +599,7 @@ def forward( attn_output = self.attn( hidden_states=norm_hidden_states, image_rotary_emb=image_rotary_emb, + offload_attention=offload_attention, **joint_attention_kwargs, ) @@ -640,6 +657,7 @@ def forward( temb_mod_params_txt: Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...], image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, joint_attention_kwargs: Optional[Dict[str, Any]] = None, + offload_attention: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: joint_attention_kwargs = joint_attention_kwargs or {} @@ -660,6 +678,7 @@ def forward( hidden_states=norm_hidden_states, encoder_hidden_states=norm_encoder_hidden_states, image_rotary_emb=image_rotary_emb, + offload_attention=offload_attention, **joint_attention_kwargs, ) @@ -877,6 +896,7 @@ class Flux2Transformer2DModel( """ _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _no_split_modules = ["Flux2TransformerBlock", "Flux2SingleTransformerBlock"] _skip_layerwise_casting_patterns = ["pos_embed", "norm"] _repeated_blocks = ["Flux2TransformerBlock", "Flux2SingleTransformerBlock"] @@ -982,6 +1002,9 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None total_layers = num_layers + num_single_layers self._musubi_block_swap = MusubiBlockSwapManager.build( depth=total_layers, @@ -1013,6 +1036,15 @@ def unfuse_qkv_projections(self): def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router, routes: List[Dict[str, Any]]): """ Set the TREAD router for efficient token routing during training. @@ -1226,9 +1258,56 @@ def forward( if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, hidden_states.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and self._tread_router is None + and hidden_states_buffer is None + and not musubi_offload_active + ) + segmented_checkpoint_fn = None + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + segmented_checkpoint_kwargs = {"use_reentrant": False} + # 4. Double Stream Transformer Blocks capture_idx = 0 + if use_segmented_checkpointing: + current_concat_rotary_emb = concat_rotary_emb + + def run_double_block(_idx, block, encoder_hidden_states, hidden_states): + return block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb_mod_params_img=double_stream_mod_img, + temb_mod_params_txt=double_stream_mod_txt, + image_rotary_emb=current_concat_rotary_emb, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + encoder_hidden_states, hidden_states = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (encoder_hidden_states, hidden_states), + run_double_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + capture_idx = num_double + for index_block, block in enumerate(self.transformer_blocks): + if use_segmented_checkpointing: + break + global_layer_idx = index_block if musubi_offload_active and musubi_manager.is_managed_block(global_layer_idx): @@ -1275,11 +1354,11 @@ def forward( def create_custom_forward(module): def custom_forward(*inputs): - return module(*inputs) + return module(*inputs, offload_attention=self.gradient_checkpointing_offload_attention) return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1304,6 +1383,7 @@ def custom_forward(*inputs): temb_mod_params_txt=double_stream_mod_txt, image_rotary_emb=current_concat_rotary_emb, joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(global_layer_idx): @@ -1328,7 +1408,33 @@ def custom_forward(*inputs): current_concat_pe = concat_rotary_emb # 5. Single Stream Transformer Blocks + if use_segmented_checkpointing: + + def run_single_block(_idx, block, hidden_states): + return block( + hidden_states=hidden_states, + encoder_hidden_states=None, + temb_mod_params=single_stream_mod, + image_rotary_emb=current_concat_pe, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + (hidden_states,) = checkpoint_sequential_state( + self.single_transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_single_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + capture_idx += len(self.single_transformer_blocks) + for index_block, block in enumerate(self.single_transformer_blocks): + if use_segmented_checkpointing: + break + global_layer_idx = num_double + index_block if musubi_offload_active and musubi_manager.is_managed_block(global_layer_idx): @@ -1382,11 +1488,11 @@ def custom_forward(*inputs): def create_custom_forward(module): def custom_forward(*inputs): - return module(*inputs) + return module(*inputs, offload_attention=self.gradient_checkpointing_offload_attention) return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1409,6 +1515,7 @@ def custom_forward(*inputs): temb_mod_params=single_stream_mod, image_rotary_emb=current_concat_pe, joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(global_layer_idx): diff --git a/simpletuner/helpers/models/heartmula/model.py b/simpletuner/helpers/models/heartmula/model.py index 7d5ffb7c6..d657a7c99 100644 --- a/simpletuner/helpers/models/heartmula/model.py +++ b/simpletuner/helpers/models/heartmula/model.py @@ -50,6 +50,7 @@ class HeartMuLa(AudioModelFoundation): def __init__(self, config, accelerator): super().__init__(config, accelerator) self.model = None + self.vae = None self.tokenizer: Optional[Tokenizer] = None self.gen_config: Optional[HeartMuLaGenConfig] = None self._gen_asset_path: Optional[str] = None @@ -97,6 +98,10 @@ def _validate_audio_dataset_usage(config: dict) -> list[ValidationResult]: ] return [] + def check_user_config(self): + super().check_user_config() + self.DEFAULT_PIPELINE_TYPE = PipelineTypes.TEXT2AUDIO + def uses_audio_latents(self) -> bool: return False @@ -173,6 +178,26 @@ def load_model(self, move_to_device: bool = True): self.model.to(self.accelerator.device) return self.model + def get_pipeline(self, pipeline_type: str = PipelineTypes.TEXT2AUDIO, load_base_model: bool = True): + if isinstance(pipeline_type, str): + pipeline_type = PipelineTypes(pipeline_type) + if pipeline_type != PipelineTypes.TEXT2AUDIO: + raise NotImplementedError(f"Pipeline type {pipeline_type} not defined in {self.__class__.__name__}.") + cached_pipeline = self.pipelines.get(pipeline_type) + transformer = self.unwrap_model(model=self.model) if self.model is not None else None + if cached_pipeline is not None: + if transformer is not None: + cached_pipeline.transformer = transformer + return cached_pipeline + + pipeline = self.PIPELINE_CLASSES[pipeline_type].from_pretrained( + self._resolve_pretrained_path(), + transformer=transformer, + torch_dtype=self.config.weight_dtype, + ) + self.pipelines[pipeline_type] = pipeline + return pipeline + def add_lora_adapter(self): from peft import LoraConfig, get_peft_model @@ -222,6 +247,23 @@ def add_lora_adapter(self): return addkeys, misskeys + def save_lora_weights(self, *args, **kwargs): + if args: + save_directory = args[0] + else: + save_directory = kwargs.get("save_directory") + if save_directory is None: + raise ValueError("save_directory is required to save LoRA weights.") + + os.makedirs(save_directory, exist_ok=True) + trained_component = self.get_trained_component() + if trained_component is None: + raise ValueError("No HeartMuLa component is available to save LoRA weights.") + trained_component = self.unwrap_model(model=trained_component) + if not hasattr(trained_component, "save_pretrained"): + raise NotImplementedError("HeartMuLa LoRA saving requires a PEFT-wrapped component.") + trained_component.save_pretrained(save_directory, safe_serialization=True) + def prepare_batch(self, batch: dict, state: dict) -> dict: if not batch: return batch diff --git a/simpletuner/helpers/models/hidream/model.py b/simpletuner/helpers/models/hidream/model.py index d0d0f716d..17b583f1c 100644 --- a/simpletuner/helpers/models/hidream/model.py +++ b/simpletuner/helpers/models/hidream/model.py @@ -762,6 +762,25 @@ def get_lora_target_layers(self): ) return targets + @staticmethod + def _precision_backend(precision: Optional[str]) -> Optional[str]: + if not isinstance(precision, str) or precision in ("", "no_change"): + return None + precision = precision.lower() + if "quanto" in precision: + return "quanto" + if "torchao" in precision: + return "torchao" + if "sdnq" in precision: + return "sdnq" + if "bnb" in precision: + return "bnb" + if precision == "fp8-native": + return "fp8-native" + if precision == "fp8-transformerengine": + return "fp8-transformerengine" + return None + def check_user_config(self): """ Checks self.config values against important issues. Optionally implemented in child class. @@ -770,6 +789,23 @@ def check_user_config(self): raise ValueError( f"{self.NAME} does not support fp8-quanto. Please use fp8-torchao or int8 precision level instead." ) + base_backend = self._precision_backend(self.config.base_model_precision) + text_encoder_4_precision = getattr(self.config, "text_encoder_4_precision", None) + text_encoder_4_backend = self._precision_backend(text_encoder_4_precision) + if base_backend and text_encoder_4_backend and base_backend != text_encoder_4_backend: + if text_encoder_4_precision == "int4-quanto": + logger.warning( + "HiDream's bundled Llama text encoder int4-quanto setting is only compatible with Quanto " + "base quantisation; setting text_encoder_4_precision=no_change for %s.", + self.config.base_model_precision, + ) + self.config.text_encoder_4_precision = "no_change" + else: + raise ValueError( + f"{self.NAME} cannot mix base model precision {self.config.base_model_precision!r} with " + f"text_encoder_4_precision={text_encoder_4_precision!r}. Use one quant backend or set " + "text_encoder_4_precision=no_change." + ) t5_max_length = 128 if self.config.tokenizer_max_length is None or self.config.tokenizer_max_length == 0: logger.warning(f"Setting T5 XXL tokeniser max length to {t5_max_length} for {self.NAME}.") diff --git a/simpletuner/helpers/models/hidream/transformer.py b/simpletuner/helpers/models/hidream/transformer.py index 582868298..11b99832e 100644 --- a/simpletuner/helpers/models/hidream/transformer.py +++ b/simpletuner/helpers/models/hidream/transformer.py @@ -23,6 +23,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter @@ -1091,6 +1092,8 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None total_layers = num_layers + num_single_layers self._musubi_block_swap = MusubiBlockSwapManager.build( @@ -1112,7 +1115,13 @@ def set_router(self, router: TREADRouter, routes: Optional[List[Dict]] = None): def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend - def _set_gradient_checkpointing(self, module, value=False): + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + + def _set_gradient_checkpointing(self, module=None, value=False, enable=None, gradient_checkpointing_func=None): """ Recursively enables or disables gradient checkpointing for all modules. @@ -1120,6 +1129,13 @@ def _set_gradient_checkpointing(self, module, value=False): module: Module to set gradient checkpointing for value: Whether to enable (True) or disable (False) gradient checkpointing """ + if enable is not None: + value = enable + if gradient_checkpointing_func is not None: + self._gradient_checkpointing_func = gradient_checkpointing_func + if module is None: + module = self + self.gradient_checkpointing = value if isinstance( module, ( @@ -1134,7 +1150,7 @@ def _set_gradient_checkpointing(self, module, value=False): # Also set checkpointing for child modules that might not be directly accessible for child in module.children(): - self._set_gradient_checkpointing(child, value) + self._set_gradient_checkpointing(child, value=value) def enable_gradient_checkpointing(self): """Enables gradient checkpointing for the model""" @@ -1304,7 +1320,7 @@ def custom_forward(proj): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1423,7 +1439,7 @@ def custom_forward(p_embedder): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1482,7 +1498,7 @@ def custom_forward(embedder): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1527,7 +1543,7 @@ def custom_forward(proj): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1584,7 +1600,7 @@ def custom_forward(pe_embedder): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1682,7 +1698,12 @@ def custom_forward(pe_embedder): ) # Process through the block with optional gradient checkpointing - if self.training and self.gradient_checkpointing: + if self.training and should_checkpoint_block( + bid, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): def create_custom_forward(module, return_dict=None): def custom_forward(*inputs): @@ -1693,7 +1714,7 @@ def custom_forward(*inputs): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1814,7 +1835,12 @@ def custom_forward(*inputs): hidden_states = torch.cat([hidden_states, cur_llama_embedding], dim=1) # Process through the block with optional gradient checkpointing - if self.training and self.gradient_checkpointing: + if self.training and should_checkpoint_block( + len(self.double_stream_blocks) + bid, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): def create_custom_forward(module, return_dict=None): def custom_forward(*inputs): @@ -1825,7 +1851,7 @@ def custom_forward(*inputs): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1914,7 +1940,7 @@ def custom_forward(final_layer): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/hunyuanvideo/autoencoder.py b/simpletuner/helpers/models/hunyuanvideo/autoencoder.py index d7723abe7..35d89ad0d 100644 --- a/simpletuner/helpers/models/hunyuanvideo/autoencoder.py +++ b/simpletuner/helpers/models/hunyuanvideo/autoencoder.py @@ -1095,10 +1095,18 @@ def set_tile_sample_min_size(self, sample_size: int, tile_overlap_factor: float self.tile_latent_min_size * self.tile_overlap_factor ).is_integer(), "self.tile_latent_min_size multiplied by tile_overlap_factor must be an integer" - def _set_gradient_checkpointing(self, module, value=False): + def _set_gradient_checkpointing(self, module=None, value=False, enable=None, gradient_checkpointing_func=None): """Enable or disable gradient checkpointing on encoder and decoder.""" + if enable is not None: + value = enable + if gradient_checkpointing_func is not None: + self._gradient_checkpointing_func = gradient_checkpointing_func + if module is None: + module = self if isinstance(module, (Encoder, Decoder)): module.gradient_checkpointing = value + for child in module.children(): + self._set_gradient_checkpointing(child, value=value) def enable_temporal_tiling(self, use_tiling: bool = True): raise RuntimeError("Temporal tiling is not supported for this VAE.") diff --git a/simpletuner/helpers/models/hunyuanvideo/model.py b/simpletuner/helpers/models/hunyuanvideo/model.py index 524059ec1..53396447c 100644 --- a/simpletuner/helpers/models/hunyuanvideo/model.py +++ b/simpletuner/helpers/models/hunyuanvideo/model.py @@ -75,6 +75,7 @@ class HunyuanVideo(VideoModelFoundation): "i2v-480p-distilled": "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_i2v_distilled", "i2v-720p-distilled": "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_i2v_distilled", } + UPSTREAM_MODEL_REPO = "tencent/HunyuanVideo-1.5" MODEL_LICENSE = "agpl-3.0" # Component repositories - direct loading without subfolders @@ -256,6 +257,16 @@ def _is_i2v_like_flavour(self) -> bool: except Exception: return False + def check_user_config(self): + super().check_user_config() + flavour = getattr(self.config, "model_flavour", None) or self.DEFAULT_MODEL_FLAVOUR + if flavour not in self.HUGGINGFACE_PATHS: + raise ValueError(f"Unsupported HunyuanVideo flavour '{flavour}'. Expected one of {list(self.HUGGINGFACE_PATHS)}") + + model_path = getattr(self.config, "pretrained_model_name_or_path", None) + if model_path in (None, "", "None", self.UPSTREAM_MODEL_REPO): + self.config.pretrained_model_name_or_path = self.HUGGINGFACE_PATHS[flavour] + def requires_conditioning_dataset(self) -> bool: return self._is_i2v_like_flavour() or super().requires_conditioning_dataset() @@ -282,7 +293,11 @@ def load_text_encoder(self, move_to_device: bool = True): Load the Qwen2.5 VL text encoder and ByT5 glyph encoder. """ device = self.accelerator.device if move_to_device else torch.device("cpu") - qwen_path = getattr(self.config, "hunyuan_text_encoder_path", None) or self.TEXT_ENCODER_REPO + qwen_path = ( + self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + or getattr(self.config, "hunyuan_text_encoder_path", None) + or self.TEXT_ENCODER_REPO + ) logger.info(f"Loading HunyuanVideo text encoder from {qwen_path}") tokenizer = Qwen2Tokenizer.from_pretrained(qwen_path) diff --git a/simpletuner/helpers/models/hunyuanvideo/transformer.py b/simpletuner/helpers/models/hunyuanvideo/transformer.py index 6640c477f..35c5debdb 100644 --- a/simpletuner/helpers/models/hunyuanvideo/transformer.py +++ b/simpletuner/helpers/models/hunyuanvideo/transformer.py @@ -45,6 +45,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -168,11 +170,13 @@ def __call__( value = torch.cat([value, encoder_value], dim=1) batch_size, seq_len, heads, dim = query.shape - attention_mask = F.pad(attention_mask, (seq_len - attention_mask.shape[1], 0), value=True) - attention_mask = attention_mask.bool() - self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1) - self_attn_mask_2 = self_attn_mask_1.transpose(2, 3) - attention_mask = (self_attn_mask_1 & self_attn_mask_2).bool() + if attention_mask is not None: + attention_mask = F.pad(attention_mask, (seq_len - attention_mask.shape[1], 0), value=True) + attention_mask = attention_mask.bool() + if bool(attention_mask.all()): + attention_mask = None + else: + attention_mask = attention_mask.view(batch_size, 1, 1, seq_len) # 5. Attention hidden_states = dispatch_attention_fn( @@ -645,6 +649,9 @@ def forward( context_temb: Optional[torch.Tensor] = None, attention_mask: Optional[torch.Tensor] = None, freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, *args, **kwargs, ) -> Tuple[torch.Tensor, torch.Tensor]: @@ -659,12 +666,13 @@ def forward( ) # 2. Joint attention - attn_output, context_attn_output = self.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - attention_mask=attention_mask, - image_rotary_emb=freqs_cis, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output, context_attn_output = self.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + attention_mask=attention_mask, + image_rotary_emb=freqs_cis, + ) # 3. Modulation and residual connection if gate_msa.ndim == 2: @@ -687,8 +695,17 @@ def forward( norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp) + c_shift_mlp # 4. Feed-forward - ff_output = self.ff(norm_hidden_states) - context_ff_output = self.ff_context(norm_encoder_hidden_states) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + ff_output, context_ff_output = checkpoint_fn( + self._ffn_forward, + norm_hidden_states, + norm_encoder_hidden_states, + use_reentrant=False, + ) + else: + ff_output, context_ff_output = self._ffn_forward(norm_hidden_states, norm_encoder_hidden_states) if gate_mlp.ndim == 2: gate_mlp = gate_mlp.unsqueeze(1) @@ -699,6 +716,13 @@ def forward( return hidden_states, encoder_hidden_states + def _ffn_forward( + self, + norm_hidden_states: torch.Tensor, + norm_encoder_hidden_states: torch.Tensor, + ) -> Tuple[torch.Tensor, torch.Tensor]: + return self.ff(norm_hidden_states), self.ff_context(norm_encoder_hidden_states) + class HunyuanVideo15Transformer3DModel( ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin, AttentionMixin @@ -740,6 +764,7 @@ class HunyuanVideo15Transformer3DModel( """ _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _skip_layerwise_casting_patterns = ["x_embedder", "context_embedder", "norm"] _no_split_modules = [ "HunyuanVideo15TransformerBlock", @@ -795,6 +820,7 @@ def __init__( self._tread_router = None self._tread_routes = None self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None inner_dim = num_attention_heads * attention_head_dim out_channels = out_channels or in_channels @@ -839,6 +865,7 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, blocks_to_swap=musubi_blocks_to_swap, @@ -1002,6 +1029,11 @@ def forward( encoder_hidden_states = torch.stack(new_encoder_hidden_states) encoder_attention_mask = torch.stack(new_encoder_attention_mask) + live_conditioning_columns = encoder_attention_mask.any(dim=0) + if bool(live_conditioning_columns.any()): + last_live_conditioning_column = int(live_conditioning_columns.nonzero()[-1].item()) + 1 + encoder_hidden_states = encoder_hidden_states[:, :last_live_conditioning_column] + encoder_attention_mask = encoder_attention_mask[:, :last_live_conditioning_column] grad_enabled = torch.is_grad_enabled() musubi_manager = self._musubi_block_swap @@ -1012,37 +1044,92 @@ def forward( # 4. Transformer blocks capture_idx = 0 if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = torch.utils.checkpoint.checkpoint - for idx, block in enumerate(self.transformer_blocks): - if musubi_offload_active and musubi_manager.is_managed_block(idx): - musubi_manager.stream_in(block, hidden_states.device) - hidden_states, encoder_hidden_states = checkpoint_fn( - block, - hidden_states, - encoder_hidden_states, - temb_hidden, - temb_context, - encoder_attention_mask, - image_rotary_emb, - use_reentrant=False, - ) - if musubi_offload_active and musubi_manager.is_managed_block(idx): - musubi_manager.stream_out(block) - tokens_view = _reshape_hunyuan_video_tokens( - hidden_states, - batch_size=batch_size, - post_patch_num_frames=post_patch_num_frames, - post_patch_height=post_patch_height, - post_patch_width=post_patch_width, + use_sequential_segments = ( + self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and hidden_states_buffer is None + and not self.gradient_checkpointing_backend.endswith("-ffn") + ) + + if use_sequential_segments: + + def run_segment_block(idx, block, segment_hidden_states, segment_encoder_hidden_states): + if musubi_offload_active and musubi_manager.is_managed_block(idx): + musubi_manager.stream_in(block, segment_hidden_states.device) + result = block( + segment_hidden_states, + segment_encoder_hidden_states, + temb_hidden, + temb_context, + encoder_attention_mask, + image_rotary_emb, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + if musubi_offload_active and musubi_manager.is_managed_block(idx): + musubi_manager.stream_out(block) + return result + + hidden_states, encoder_hidden_states = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states, encoder_hidden_states), + run_segment_block, + checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, ) - _store_hidden_state(hidden_states_buffer, f"layer_{capture_idx}", tokens_view) - capture_idx += 1 + else: + for idx, block in enumerate(self.transformer_blocks): + if musubi_offload_active and musubi_manager.is_managed_block(idx): + musubi_manager.stream_in(block, hidden_states.device) + checkpoint_this_block = should_checkpoint_block( + idx, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + if checkpoint_this_block and not self.gradient_checkpointing_backend.endswith("-ffn"): + hidden_states, encoder_hidden_states = checkpoint_fn( + block, + hidden_states, + encoder_hidden_states, + temb_hidden, + temb_context, + encoder_attention_mask, + image_rotary_emb, + offload_attention=self.gradient_checkpointing_offload_attention, + use_reentrant=False, + ) + else: + hidden_states, encoder_hidden_states = block( + hidden_states, + encoder_hidden_states, + temb_hidden, + temb_context, + encoder_attention_mask, + image_rotary_emb, + checkpoint_ffn=self.gradient_checkpointing_backend.endswith("-ffn") and checkpoint_this_block, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + if musubi_offload_active and musubi_manager.is_managed_block(idx): + musubi_manager.stream_out(block) + tokens_view = _reshape_hunyuan_video_tokens( + hidden_states, + batch_size=batch_size, + post_patch_num_frames=post_patch_num_frames, + post_patch_height=post_patch_height, + post_patch_width=post_patch_width, + ) + _store_hidden_state(hidden_states_buffer, f"layer_{capture_idx}", tokens_view) + capture_idx += 1 else: for idx, block in enumerate(self.transformer_blocks): @@ -1055,6 +1142,7 @@ def forward( temb_context, encoder_attention_mask, image_rotary_emb, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(idx): musubi_manager.stream_out(block) @@ -1096,5 +1184,11 @@ def set_router(self, router, routes=None): def set_gradient_checkpointing_interval(self, interval: int): self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) diff --git a/simpletuner/helpers/models/ideogram/model.py b/simpletuner/helpers/models/ideogram/model.py index b3407f41d..9d585a711 100644 --- a/simpletuner/helpers/models/ideogram/model.py +++ b/simpletuner/helpers/models/ideogram/model.py @@ -103,12 +103,13 @@ def load_vae(self, move_to_device: bool = True): def load_text_encoder(self, move_to_device: bool = True): repo_id = getattr(self.config, "pretrained_model_name_or_path", None) or self.HUGGINGFACE_PATHS["fp8"] pipe_config = Ideogram4PipelineConfig(weights_repo=repo_id) + qwen_repo_id = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") tokenizer, text_encoder = _load_qwen3_vl( - repo_id, + qwen_repo_id or repo_id, self.accelerator.device, self.config.weight_dtype, - tokenizer_subfolder=pipe_config.tokenizer_subfolder, - text_encoder_subfolder=pipe_config.text_encoder_subfolder, + tokenizer_subfolder=None if qwen_repo_id else pipe_config.tokenizer_subfolder, + text_encoder_subfolder=None if qwen_repo_id else pipe_config.text_encoder_subfolder, ) self.tokenizers = [tokenizer] self.text_encoders = [text_encoder] diff --git a/simpletuner/helpers/models/ideogram/transformer.py b/simpletuner/helpers/models/ideogram/transformer.py index de236f427..c30675f41 100644 --- a/simpletuner/helpers/models/ideogram/transformer.py +++ b/simpletuner/helpers/models/ideogram/transformer.py @@ -29,6 +29,7 @@ QWEN3_VL_ACTIVATION_LAYERS, ) from simpletuner.helpers.models.ideogram.quantized_loading import Fp8Linear +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block @dataclass @@ -336,6 +337,8 @@ def __init__(self, config: Ideogram4Config) -> None: ) self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None def enable_gradient_checkpointing(self) -> None: self.gradient_checkpointing = True @@ -346,6 +349,12 @@ def disable_gradient_checkpointing(self) -> None: def set_gradient_checkpointing_backend(self, backend: str) -> None: self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int) -> None: + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None) -> None: + self.gradient_checkpointing_segment_stride = segment_stride + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.flowmap_deltatime_type = validate_flowmap_deltatime_type(deltatime_type, model_name="Ideogram") if self.delta_t_embedding is None: @@ -461,15 +470,23 @@ def forward( sin = sin.to(h.dtype) if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = torch.utils.checkpoint.checkpoint - for layer in self.layers: - h = checkpoint_fn(layer, h, segment_ids, cos, sin, adaln_input, use_reentrant=False) + for layer_idx, layer in enumerate(self.layers): + if should_checkpoint_block( + layer_idx, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + h = checkpoint_fn(layer, h, segment_ids, cos, sin, adaln_input, use_reentrant=False) + else: + h = layer(h, segment_ids=segment_ids, cos=cos, sin=sin, adaln_input=adaln_input) else: for layer in self.layers: h = layer(h, segment_ids=segment_ids, cos=cos, sin=sin, adaln_input=adaln_input) diff --git a/simpletuner/helpers/models/kandinsky5_image/pipeline_kandinsky5_t2i.py b/simpletuner/helpers/models/kandinsky5_image/pipeline_kandinsky5_t2i.py index 7eb180a52..1ab708a10 100644 --- a/simpletuner/helpers/models/kandinsky5_image/pipeline_kandinsky5_t2i.py +++ b/simpletuner/helpers/models/kandinsky5_image/pipeline_kandinsky5_t2i.py @@ -426,10 +426,10 @@ def __call__( image=kwargs.get("image"), ) - text_rope_pos = torch.arange(prompt_cu_seqlens.diff().max().item(), device=device) + text_rope_pos = torch.arange(prompt_embeds_qwen.shape[1], device=device) negative_text_rope_pos = ( - torch.arange(negative_prompt_cu_seqlens.diff().max().item(), device=device) - if negative_prompt_cu_seqlens is not None + torch.arange(negative_prompt_embeds_qwen.shape[1], device=device) + if negative_prompt_embeds_qwen is not None else None ) @@ -534,28 +534,13 @@ def __call__( if self.vae is None: raise ValueError("VAE is not loaded; set output_type='latent' or load the VAE to decode images.") latents = latents.to(self.vae.dtype) - images = latents.reshape( - batch_size, - num_images_per_prompt, - 1, - height // self.vae_scale_factor_spatial, - width // self.vae_scale_factor_spatial, - num_channels_latents, - ) - images = images.permute(0, 1, 5, 2, 3, 4) - images = images.reshape( - batch_size * num_images_per_prompt, - num_channels_latents, - 1, - height // self.vae_scale_factor_spatial, - width // self.vae_scale_factor_spatial, - ) + images = latents[:, 0].permute(0, 3, 1, 2).contiguous() images = images / getattr(self.vae.config, "scaling_factor", 1.0) images = self.vae.decode(images).sample if self.image_processor is None: images = images else: - images = self.image_processor.postprocess_image(images, output_type=output_type) + images = self.image_processor.postprocess(images, output_type=output_type) else: images = latents diff --git a/simpletuner/helpers/models/kandinsky5_video/pipeline_kandinsky5_t2v.py b/simpletuner/helpers/models/kandinsky5_video/pipeline_kandinsky5_t2v.py index 7fa5633e9..c02050062 100644 --- a/simpletuner/helpers/models/kandinsky5_video/pipeline_kandinsky5_t2v.py +++ b/simpletuner/helpers/models/kandinsky5_video/pipeline_kandinsky5_t2v.py @@ -785,11 +785,11 @@ def __call__( torch.arange(width // self.vae_scale_factor_spatial // 2, device=device), ] - text_rope_pos = torch.arange(prompt_cu_seqlens.diff().max().item(), device=device) + text_rope_pos = torch.arange(prompt_embeds_qwen.shape[1], device=device) negative_text_rope_pos = ( - torch.arange(negative_prompt_cu_seqlens.diff().max().item(), device=device) - if negative_prompt_cu_seqlens is not None + torch.arange(negative_prompt_embeds_qwen.shape[1], device=device) + if negative_prompt_embeds_qwen is not None else None ) diff --git a/simpletuner/helpers/models/kandinsky5_video/transformer_kandinsky5.py b/simpletuner/helpers/models/kandinsky5_video/transformer_kandinsky5.py index 82770db97..094fef576 100644 --- a/simpletuner/helpers/models/kandinsky5_video/transformer_kandinsky5.py +++ b/simpletuner/helpers/models/kandinsky5_video/transformer_kandinsky5.py @@ -39,7 +39,9 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter @@ -407,7 +409,15 @@ def __init__(self): if not hasattr(F, "scaled_dot_product_attention"): raise ImportError(f"{self.__class__.__name__} requires PyTorch 2.0. Please upgrade your pytorch version.") - def __call__(self, attn, hidden_states, encoder_hidden_states=None, rotary_emb=None, sparse_params=None): + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + rotary_emb=None, + sparse_params=None, + offload_attention: bool = False, + ): def _describe_tensor(name: str, tensor: Optional[Tensor]) -> str: if tensor is None: return f"{name}=None" @@ -474,13 +484,14 @@ def apply_rotary(x, rope): getattr(attn, "to_query", None) and attn.to_query.weight, getattr(attn, "to_key", None) and attn.to_key.weight, ) - attn_output = dispatch_attention_fn( - query, - key, - value, - attn_mask=attn_mask, - backend=self._attention_backend, - ) + with activation_offload_context(offload_attention, label=f"{attn.__class__.__qualname__}:attention"): + attn_output = dispatch_attention_fn( + query, + key, + value, + attn_mask=attn_mask, + backend=self._attention_backend, + ) except Exception: logger.error( "dispatch_attention_fn failed (backend=%s): %s; %s; %s; hidden_states=%s; encoder_hidden_states=%s; attn_mask=%s; %s", @@ -529,6 +540,7 @@ def forward( encoder_hidden_states: Optional[torch.Tensor] = None, sparse_params: Optional[torch.Tensor] = None, rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + offload_attention: bool = False, **kwargs, ) -> torch.Tensor: attn_parameters = set(inspect.signature(self.processor.__call__).parameters.keys()) @@ -546,6 +558,7 @@ def forward( encoder_hidden_states=encoder_hidden_states, sparse_params=sparse_params, rotary_emb=rotary_emb, + offload_attention=offload_attention, **kwargs, ) @@ -607,7 +620,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - def forward(self, x, time_embed, rope): + def forward(self, x, time_embed, rope, offload_attention: bool = False): self_attn_params, ff_params = torch.chunk( _broadcast_modulation_params(self.text_modulation(time_embed), x), 2, dim=-1 ) @@ -627,7 +640,7 @@ def forward(self, x, time_embed, rope): ] ) raise RuntimeError(debug_msg) from err - out = self.self_attention(out, rotary_emb=rope) + out = self.self_attention(out, rotary_emb=rope, offload_attention=offload_attention) x = (x.float() + gate.float() * out.float()).type_as(x) shift, scale, gate = torch.chunk(ff_params, 3, dim=-1) @@ -652,7 +665,7 @@ def __init__(self, model_dim, time_dim, ff_dim, head_dim): self.feed_forward_norm = nn.LayerNorm(model_dim, elementwise_affine=False) self.feed_forward = Kandinsky5FeedForward(model_dim, ff_dim) - def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params): + def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params, offload_attention: bool = False): self_attn_params, cross_attn_params, ff_params = torch.chunk( _broadcast_modulation_params(self.visual_modulation(time_embed), visual_embed), 3, dim=-1 ) @@ -661,14 +674,23 @@ def forward(self, visual_embed, text_embed, time_embed, rope, sparse_params): visual_out = (self.self_attention_norm(visual_embed.float()) * (scale.float() + 1.0) + shift.float()).type_as( visual_embed ) - visual_out = self.self_attention(visual_out, rotary_emb=rope, sparse_params=sparse_params) + visual_out = self.self_attention( + visual_out, + rotary_emb=rope, + sparse_params=sparse_params, + offload_attention=offload_attention, + ) visual_embed = (visual_embed.float() + gate.float() * visual_out.float()).type_as(visual_embed) shift, scale, gate = torch.chunk(cross_attn_params, 3, dim=-1) visual_out = (self.cross_attention_norm(visual_embed.float()) * (scale.float() + 1.0) + shift.float()).type_as( visual_embed ) - visual_out = self.cross_attention(visual_out, encoder_hidden_states=text_embed) + visual_out = self.cross_attention( + visual_out, + encoder_hidden_states=text_embed, + offload_attention=offload_attention, + ) visual_embed = (visual_embed.float() + gate.float() * visual_out.float()).type_as(visual_embed) shift, scale, gate = torch.chunk(ff_params, 3, dim=-1) @@ -701,6 +723,7 @@ class Kandinsky5Transformer3DModel( "Kandinsky5TransformerDecoderBlock", ] _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _cp_plan = { "": { "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), @@ -784,6 +807,9 @@ def __init__( self.out_layer = Kandinsky5OutLayer(model_dim, time_dim, out_visual_dim, patch_size) self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_text_blocks + num_visual_blocks, blocks_to_swap=musubi_blocks_to_swap, @@ -798,6 +824,15 @@ def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_t def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): """Attach a TREAD router and route definitions.""" self._tread_router = router @@ -891,20 +926,38 @@ def forward( f"Tokenwise timestep count {visual_time_embed.shape[1]} does not match visual token count {expected_tokens}." ) - for text_transformer_block in self.text_transformer_blocks: - if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + for text_layer_idx, text_transformer_block in enumerate(self.text_transformer_blocks): + if torch.is_grad_enabled() and should_checkpoint_block( + text_layer_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = torch.utils.checkpoint.checkpoint + def _checkpointed_text_block(text_embed, text_time_embed, text_rope, block=text_transformer_block): + return block( + text_embed, + text_time_embed, + text_rope, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + text_embed = checkpoint_fn( - text_transformer_block, text_embed, text_time_embed, text_rope, use_reentrant=False + _checkpointed_text_block, text_embed, text_time_embed, text_rope, use_reentrant=False ) else: - text_embed = text_transformer_block(text_embed, text_time_embed, text_rope) + text_embed = text_transformer_block( + text_embed, + text_time_embed, + text_rope, + offload_attention=self.gradient_checkpointing_offload_attention, + ) visual_rope = self.visual_rope_embeddings(visual_shape, visual_rope_pos, scale_factor) to_fractal = sparse_params["to_fractal"] if sparse_params is not None else False @@ -982,16 +1035,38 @@ def _to_pos(idx: int) -> int: current_rope = self._route_rope(visual_rope, tread_mask_info, keep_len=visual_embed.size(1)) routing_now = True - if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if torch.is_grad_enabled() and should_checkpoint_block( + global_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = torch.utils.checkpoint.checkpoint + def _checkpointed_visual_block( + visual_embed, + text_embed, + visual_time_embed, + current_rope, + sparse_params, + block=visual_transformer_block, + ): + return block( + visual_embed, + text_embed, + visual_time_embed, + current_rope, + sparse_params, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + visual_embed = checkpoint_fn( - visual_transformer_block, + _checkpointed_visual_block, visual_embed, text_embed, visual_time_embed, @@ -1001,7 +1076,12 @@ def _to_pos(idx: int) -> int: ) else: visual_embed = visual_transformer_block( - visual_embed, text_embed, visual_time_embed, current_rope, sparse_params + visual_embed, + text_embed, + visual_time_embed, + current_rope, + sparse_params, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(visual_transformer_block, "fuser"): diff --git a/simpletuner/helpers/models/kolors/controlnet.py b/simpletuner/helpers/models/kolors/controlnet.py index c6870e50f..c40c72819 100644 --- a/simpletuner/helpers/models/kolors/controlnet.py +++ b/simpletuner/helpers/models/kolors/controlnet.py @@ -720,9 +720,19 @@ def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[i for module in self.children(): fn_recursive_set_attention_slice(module, reversed_slice_size) - def _set_gradient_checkpointing(self, module, value: bool = False) -> None: + def _set_gradient_checkpointing( + self, module=None, value: bool = False, enable=None, gradient_checkpointing_func=None + ) -> None: + if enable is not None: + value = enable + if gradient_checkpointing_func is not None: + self._gradient_checkpointing_func = gradient_checkpointing_func + if module is None: + module = self if isinstance(module, (CrossAttnDownBlock2D, DownBlock2D)): module.gradient_checkpointing = value + for child in module.children(): + self._set_gradient_checkpointing(child, value=value) def forward( self, diff --git a/simpletuner/helpers/models/krea2/model.py b/simpletuner/helpers/models/krea2/model.py index 4d4e803b9..5d5c86b27 100644 --- a/simpletuner/helpers/models/krea2/model.py +++ b/simpletuner/helpers/models/krea2/model.py @@ -178,8 +178,8 @@ def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self.PROCESSOR_PATH - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self.PROCESSOR_PATH) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_revision = getattr(self.config, "processor_revision", getattr(self.config, "revision", None)) processor_kwargs = {"pretrained_model_name_or_path": processor_path} diff --git a/simpletuner/helpers/models/krea2/transformer.py b/simpletuner/helpers/models/krea2/transformer.py index 177222919..b6499a81a 100644 --- a/simpletuner/helpers/models/krea2/transformer.py +++ b/simpletuner/helpers/models/krea2/transformer.py @@ -39,6 +39,9 @@ ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.attention_backend import maybe_metal_flash_rope_attention +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.tread import TREADRouter logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -229,8 +232,14 @@ def __init__(self, dim: int, num_heads: int, num_kv_heads: int, intermediate_siz self.attn = Krea2Attention(dim, num_heads, num_kv_heads, eps=eps) self.ff = Krea2SwiGLU(dim, intermediate_size) - def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None) -> torch.Tensor: - hidden_states = hidden_states + self.attn(self.norm1(hidden_states), attention_mask=attention_mask) + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + offload_attention: bool = False, + ) -> torch.Tensor: + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + hidden_states = hidden_states + self.attn(self.norm1(hidden_states), attention_mask=attention_mask) hidden_states = hidden_states + self.ff(self.norm2(hidden_states)) return hidden_states @@ -287,26 +296,49 @@ def __init__(self, hidden_size: int, intermediate_size: int, num_heads: int, num self.attn = Krea2Attention(hidden_size, num_heads, num_kv_heads, eps=norm_eps) self.ff = Krea2SwiGLU(hidden_size, intermediate_size) + def _ffn_forward( + self, + hidden_states: torch.Tensor, + postscale: torch.Tensor, + postshift: torch.Tensor, + postgate: torch.Tensor, + ) -> torch.Tensor: + ff_out = self.ff((1.0 + postscale) * self.norm2(hidden_states) + postshift) + return hidden_states + postgate * ff_out + def forward( self, hidden_states: torch.Tensor, temb: torch.Tensor, image_rotary_emb: tuple[torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None = None, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> torch.Tensor: # temb: (B, 1, 6 * hidden_size), shared across all blocks; each block only learns an additive table. modulation = temb.unflatten(-1, (6, -1)) + self.scale_shift_table prescale, preshift, pregate, postscale, postshift, postgate = modulation.unbind(-2) - attn_out = self.attn( - (1.0 + prescale) * self.norm1(hidden_states) + preshift, - attention_mask=attention_mask, - image_rotary_emb=image_rotary_emb, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_out = self.attn( + (1.0 + prescale) * self.norm1(hidden_states) + preshift, + attention_mask=attention_mask, + image_rotary_emb=image_rotary_emb, + ) hidden_states = hidden_states + pregate * attn_out - ff_out = self.ff((1.0 + postscale) * self.norm2(hidden_states) + postshift) - hidden_states = hidden_states + postgate * ff_out - return hidden_states + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + return checkpoint_fn( + self._ffn_forward, + hidden_states, + postscale, + postshift, + postgate, + use_reentrant=False, + ) + return self._ffn_forward(hidden_states, postscale, postshift, postgate) class Krea2TimestepEmbedding(nn.Module): @@ -494,6 +526,8 @@ class Krea2Transformer2DModel(ModelMixin, ConfigMixin, AttentionMixin, PeftAdapt """ _supports_gradient_checkpointing = True + _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _no_split_modules = ["Krea2TransformerBlock", "Krea2TextFusionBlock", "Krea2FinalLayer"] _repeated_blocks = ["Krea2TransformerBlock"] _keep_in_fp32_modules = ["norm", "norm1", "norm2", "norm_q", "norm_k"] @@ -535,6 +569,10 @@ def __init__( self.out_channels = in_channels self.hidden_size = hidden_size self.gradient_checkpointing = False + self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.img_in = nn.Linear(in_channels, hidden_size, bias=True) self.time_embed = Krea2TimestepEmbedding( @@ -582,6 +620,18 @@ def __init__( self._tread_router: TREADRouter | None = None self._tread_routes: list[dict[str, Any]] | None = None + def set_gradient_checkpointing_backend(self, backend: str) -> None: + self.gradient_checkpointing_backend = backend + + def set_gradient_checkpointing_offload_attention(self, enabled: bool) -> None: + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int) -> None: + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None) -> None: + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: list[dict[str, Any]] | None = None) -> None: self._tread_router = router self._tread_routes = routes @@ -726,6 +776,50 @@ def _to_pos(idx): if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, hidden_states.device, torch.is_grad_enabled()) + use_segmented_checkpointing = ( + torch.is_grad_enabled() + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not use_routing + and not skip_set + and hidden_states_buffer is None + and not musubi_offload_active + ) + segmented_checkpoint_fn = None + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + + def run_segmented_block(_idx, block, hidden_states): + return block( + hidden_states, + temb_mod, + image_rotary_emb, + attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + (hidden_states,) = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_segmented_block, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + hidden_states = hidden_states[:, text_seq_len:] + output = self.final_layer(hidden_states, temb) + + if not return_dict: + return (output,) + return Transformer2DModelOutput(sample=output) + for idx, block in enumerate(self.transformer_blocks): if musubi_offload_active and musubi_manager.is_managed_block(idx): musubi_manager.stream_in(block, hidden_states.device) @@ -754,12 +848,47 @@ def _to_pos(idx): if idx in skip_set: block_output = hidden_states else: - block_output = self._gradient_checkpointing_func( - block, hidden_states, temb_mod, image_rotary_emb, attention_mask - ) + if self.gradient_checkpointing_backend.endswith("-ffn"): + block_output = block( + hidden_states, + temb_mod, + image_rotary_emb, + attention_mask, + checkpoint_ffn=True, + checkpoint_fn=self._gradient_checkpointing_func, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_checkpointed_block( + checkpoint_hidden_states, + checkpoint_temb, + checkpoint_rope, + checkpoint_mask, + checkpoint_block=block, + ): + return checkpoint_block( + checkpoint_hidden_states, + checkpoint_temb, + checkpoint_rope, + checkpoint_mask, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + block_output = self._gradient_checkpointing_func( + run_checkpointed_block, hidden_states, temb_mod, image_rotary_emb, attention_mask + ) else: block_output = ( - hidden_states if idx in skip_set else block(hidden_states, temb_mod, image_rotary_emb, attention_mask) + hidden_states + if idx in skip_set + else block( + hidden_states, + temb_mod, + image_rotary_emb, + attention_mask, + offload_attention=self.gradient_checkpointing_offload_attention, + ) ) hidden_states = block_output diff --git a/simpletuner/helpers/models/longcat_image/model.py b/simpletuner/helpers/models/longcat_image/model.py index 86eb54f04..41671375d 100644 --- a/simpletuner/helpers/models/longcat_image/model.py +++ b/simpletuner/helpers/models/longcat_image/model.py @@ -211,8 +211,15 @@ def _load_text_processor_for_pipeline(self): text_processor = getattr(self, "text_processor", None) if text_processor is not None: return text_processor - model_path = get_model_config_path(self.config.model_family, self.config.pretrained_model_name_or_path) - text_processor = AutoProcessor.from_pretrained(model_path, subfolder="text_processor") + qwen_text_encoder_path = self._get_optional_config_model_path("qwen_text_encoder_model_name_or_path") + model_path = qwen_text_encoder_path or get_model_config_path( + self.config.model_family, + self.config.pretrained_model_name_or_path, + ) + processor_kwargs = {"pretrained_model_name_or_path": model_path} + if not qwen_text_encoder_path: + processor_kwargs["subfolder"] = "text_processor" + text_processor = AutoProcessor.from_pretrained(**processor_kwargs) self.text_processor = text_processor return text_processor diff --git a/simpletuner/helpers/models/longcat_image/pipeline.py b/simpletuner/helpers/models/longcat_image/pipeline.py index 3436fce6f..b7887d338 100644 --- a/simpletuner/helpers/models/longcat_image/pipeline.py +++ b/simpletuner/helpers/models/longcat_image/pipeline.py @@ -120,7 +120,7 @@ def __init__( tokenizer: AutoTokenizer, text_processor: AutoProcessor, transformer, - image_encoder: CLIPVisionModelWithProjection, + image_encoder: CLIPVisionModelWithProjection = None, feature_extractor: CLIPImageProcessor = None, ): super().__init__() diff --git a/simpletuner/helpers/models/longcat_image/transformer.py b/simpletuner/helpers/models/longcat_image/transformer.py index 53328b7b3..a0a23e678 100644 --- a/simpletuner/helpers/models/longcat_image/transformer.py +++ b/simpletuner/helpers/models/longcat_image/transformer.py @@ -27,6 +27,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context logger = get_logger(__name__, log_level="INFO") @@ -185,6 +187,7 @@ def _run_longcat_transformer_block( temb: torch.Tensor, context_temb: torch.Tensor, image_rotary_emb, + offload_attention: bool = False, ): norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = _longcat_apply_ada_layer_norm_zero( block.norm1, hidden_states, temb @@ -193,11 +196,12 @@ def _run_longcat_transformer_block( block.norm1_context, encoder_hidden_states, context_temb ) - attention_outputs = block.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - image_rotary_emb=image_rotary_emb, - ) + with activation_offload_context(offload_attention, label=f"{block.__class__.__qualname__}:attention"): + attention_outputs = block.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + image_rotary_emb=image_rotary_emb, + ) if len(attention_outputs) == 2: attn_output, context_attn_output = attention_outputs elif len(attention_outputs) == 3: @@ -244,6 +248,7 @@ def _run_longcat_single_transformer_block( encoder_hidden_states: torch.Tensor, temb: torch.Tensor, image_rotary_emb, + offload_attention: bool = False, ): text_seq_len = encoder_hidden_states.shape[1] hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) @@ -251,7 +256,8 @@ def _run_longcat_single_transformer_block( residual = hidden_states norm_hidden_states, gate = _longcat_apply_ada_layer_norm_zero_single(block.norm, hidden_states, temb) mlp_hidden_states = block.act_mlp(block.proj_mlp(norm_hidden_states)) - attn_output = block.attn(hidden_states=norm_hidden_states, image_rotary_emb=image_rotary_emb) + with activation_offload_context(offload_attention, label=f"{block.__class__.__qualname__}:attention"): + attn_output = block.attn(hidden_states=norm_hidden_states, image_rotary_emb=image_rotary_emb) hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) if gate.ndim == 2: @@ -271,6 +277,7 @@ class LongCatImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin): """ _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _cp_plan = { "": { "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), @@ -341,6 +348,9 @@ def __init__( self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels, bias=True) self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None + self.gradient_checkpointing_offload_attention = False self.initialize_weights() self.use_checkpoint = [True] * num_layers @@ -354,6 +364,15 @@ def __init__( logger=logger, ) + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.time_embed.enable_flowmap_time_conditioning(gate_value=gate_value, deltatime_type=deltatime_type) register_flowmap_config(self, gate_value, deltatime_type) @@ -456,7 +475,16 @@ def forward( for index_block, block in enumerate(self.transformer_blocks): if musubi_offload_active and musubi_manager.is_managed_block(capture_idx): musubi_manager.stream_in(block, hidden_states.device) - if torch.is_grad_enabled() and self.gradient_checkpointing and self.use_checkpoint[index_block]: + if ( + torch.is_grad_enabled() + and self.use_checkpoint[index_block] + and should_checkpoint_block( + capture_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + ): encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( _run_longcat_transformer_block, block, @@ -465,6 +493,7 @@ def forward( temb_img, temb_txt, image_rotary_emb, + self.gradient_checkpointing_offload_attention, ) else: encoder_hidden_states, hidden_states = _run_longcat_transformer_block( @@ -474,6 +503,7 @@ def forward( temb_img, temb_txt, image_rotary_emb, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(capture_idx): musubi_manager.stream_out(block) @@ -483,7 +513,16 @@ def forward( for index_block, block in enumerate(self.single_transformer_blocks): if musubi_offload_active and musubi_manager.is_managed_block(capture_idx): musubi_manager.stream_in(block, hidden_states.device) - if torch.is_grad_enabled() and self.gradient_checkpointing and self.use_single_checkpoint[index_block]: + if ( + torch.is_grad_enabled() + and self.use_single_checkpoint[index_block] + and should_checkpoint_block( + capture_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + ): encoder_hidden_states, hidden_states = self._gradient_checkpointing_func( _run_longcat_single_transformer_block, block, @@ -491,6 +530,7 @@ def forward( encoder_hidden_states, temb_single, image_rotary_emb, + self.gradient_checkpointing_offload_attention, ) else: encoder_hidden_states, hidden_states = _run_longcat_single_transformer_block( @@ -499,6 +539,7 @@ def forward( encoder_hidden_states, temb_single, image_rotary_emb, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(capture_idx): musubi_manager.stream_out(block) diff --git a/simpletuner/helpers/models/longcat_video/transformer.py b/simpletuner/helpers/models/longcat_video/transformer.py index 374867fcb..e15ac543f 100644 --- a/simpletuner/helpers/models/longcat_video/transformer.py +++ b/simpletuner/helpers/models/longcat_video/transformer.py @@ -23,6 +23,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context logger = logging.getLogger(__name__) _BSA_DISABLED_WARNED = False @@ -961,7 +963,17 @@ def __init__( self.ffn = FeedForwardSwiGLU(dim=hidden_size, hidden_dim=int(hidden_size * mlp_ratio)) def forward( - self, x, y, t, y_seqlen, latent_shape, num_cond_latents=None, return_kv=False, kv_cache=None, skip_crs_attn=False + self, + x, + y, + t, + y_seqlen, + latent_shape, + num_cond_latents=None, + return_kv=False, + kv_cache=None, + skip_crs_attn=False, + offload_attention=False, ): x_dtype = x.dtype @@ -987,13 +999,14 @@ def forward( x_m = modulate_fp32(self.mod_norm_attn, x.view(B, T, -1, C), shift_msa, scale_msa).view(B, N, C) - if kv_cache is not None: - attn_outputs = self.attn.forward_with_kv_cache( - x_m, shape=latent_shape, num_cond_latents=num_cond_latents, kv_cache=kv_cache - ) - kv_cache = kv_cache - else: - attn_outputs = self.attn(x_m, shape=latent_shape, num_cond_latents=num_cond_latents, return_kv=return_kv) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:self_attention"): + if kv_cache is not None: + attn_outputs = self.attn.forward_with_kv_cache( + x_m, shape=latent_shape, num_cond_latents=num_cond_latents, kv_cache=kv_cache + ) + kv_cache = kv_cache + else: + attn_outputs = self.attn(x_m, shape=latent_shape, num_cond_latents=num_cond_latents, return_kv=return_kv) if return_kv: x_s, kv_cache = attn_outputs else: @@ -1006,9 +1019,10 @@ def forward( if not skip_crs_attn: if kv_cache is not None: num_cond_latents = None - x = x + self.cross_attn( - self.pre_crs_attn_norm(x), y, y_seqlen, num_cond_latents=num_cond_latents, shape=latent_shape - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:cross_attention"): + x = x + self.cross_attn( + self.pre_crs_attn_norm(x), y, y_seqlen, num_cond_latents=num_cond_latents, shape=latent_shape + ) x = modulate_fp32(self.mod_norm_ffn, x.view(B, T, -1, C), shift_mlp, scale_mlp).view(B, N, C) @@ -1027,6 +1041,7 @@ class LongCatVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin): """ _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _cp_plan = { "": { "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), @@ -1107,6 +1122,9 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None + self.gradient_checkpointing_offload_attention = False self.text_tokens_zero_pad = text_tokens_zero_pad self._musubi_block_swap = MusubiBlockSwapManager.build( depth=depth, @@ -1118,6 +1136,15 @@ def __init__( def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.t_embedder.enable_flowmap_time_conditioning(gate_value=gate_value, deltatime_type=deltatime_type) register_flowmap_config(self, gate_value, deltatime_type) @@ -1273,8 +1300,13 @@ def _normalize_mask(mask: torch.Tensor) -> torch.Tensor: if musubi_offload_active and musubi_manager.is_managed_block(i): musubi_manager.stream_in(block, hidden_states.device) - if grad_enabled and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if grad_enabled and should_checkpoint_block( + capture_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -1292,6 +1324,7 @@ def _normalize_mask(mask: torch.Tensor) -> torch.Tensor: return_kv, kv_cache_dict.get(i, None), skip_crs_attn, + self.gradient_checkpointing_offload_attention, use_reentrant=False, ) else: @@ -1305,6 +1338,7 @@ def _normalize_mask(mask: torch.Tensor) -> torch.Tensor: return_kv, kv_cache_dict.get(i, None), skip_crs_attn, + self.gradient_checkpointing_offload_attention, ) if return_kv: diff --git a/simpletuner/helpers/models/ltxvideo/transformer.py b/simpletuner/helpers/models/ltxvideo/transformer.py index 274d247ee..7a10e4ce3 100644 --- a/simpletuner/helpers/models/ltxvideo/transformer.py +++ b/simpletuner/helpers/models/ltxvideo/transformer.py @@ -41,6 +41,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.tread import TREADRouter @@ -547,6 +549,9 @@ def __init__( nn.init.zeros_(self.time_sign_embed.weight) self.gradient_checkpointing = False + self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None # TREAD support self._tread_router = None @@ -558,6 +563,15 @@ def __init__( logger=logger, ) + def set_gradient_checkpointing_backend(self, backend: str): + self.gradient_checkpointing_backend = backend + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: Optional[List[Dict]] = None): """Set TREAD router and routes for token reduction during training.""" self._tread_router = router @@ -689,8 +703,53 @@ def forward( if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(self.transformer_blocks, hidden_states.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not musubi_offload_active + and grounding_objs is None + and hidden_states_buffer is None + and not output_hidden_states + ) + segmented_checkpoint_fn = None + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + capture_idx = 0 for bid, block in enumerate(self.transformer_blocks): + if use_segmented_checkpointing: + if bid != 0: + continue + + def run_segmented_block(_idx, segment_block, segment_hidden_states): + return segment_block( + hidden_states=segment_hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + encoder_attention_mask=encoder_attention_mask, + ) + + (hidden_states,) = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_segmented_block, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + continue + # TREAD routing for this layer if use_routing: # Check if this layer should use routing @@ -711,14 +770,30 @@ def forward( break if musubi_offload_active and musubi_manager.is_managed_block(bid): musubi_manager.stream_in(block, hidden_states.device) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( + checkpoint_this_block = should_checkpoint_block( + bid, + grad_enabled and self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + if checkpoint_this_block: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + checkpoint_fn = offloaded_checkpoint + checkpoint_kwargs = {"use_reentrant": False} + else: + checkpoint_fn = self._gradient_checkpointing_func + checkpoint_kwargs = {} + + hidden_states = checkpoint_fn( block, hidden_states, encoder_hidden_states, temb, image_rotary_emb, encoder_attention_mask, + **checkpoint_kwargs, ) else: hidden_states = block( diff --git a/simpletuner/helpers/models/ltxvideo2/model.py b/simpletuner/helpers/models/ltxvideo2/model.py index ff58b69c0..978fa9227 100644 --- a/simpletuner/helpers/models/ltxvideo2/model.py +++ b/simpletuner/helpers/models/ltxvideo2/model.py @@ -845,6 +845,13 @@ def load_model(self, move_to_device: bool = True): self.unwrap_model(model=self.model).set_gradient_checkpointing_interval( int(self.config.gradient_checkpointing_interval) ) + gradient_checkpointing_segment_stride = getattr(self.config, "gradient_checkpointing_segment_stride", None) + if gradient_checkpointing_segment_stride is not None: + if self.model is not None and hasattr(self.model, "set_gradient_checkpointing_segment_stride"): + logger.info("Setting gradient checkpointing segment stride..") + self.unwrap_model(model=self.model).set_gradient_checkpointing_segment_stride( + int(gradient_checkpointing_segment_stride) + ) self.fuse_qkv_projections() self.post_model_load_setup() diff --git a/simpletuner/helpers/models/ltxvideo2/transformer.py b/simpletuner/helpers/models/ltxvideo2/transformer.py index f666bdcce..e9aeb1afb 100644 --- a/simpletuner/helpers/models/ltxvideo2/transformer.py +++ b/simpletuner/helpers/models/ltxvideo2/transformer.py @@ -46,7 +46,9 @@ ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.packed_attention_processors import run_packed_qkv_attention from simpletuner.helpers.training.tread import TREADRouter @@ -956,6 +958,9 @@ def forward( perturbation_mask: Optional[torch.Tensor] = None, all_perturbed: Optional[bool] = None, audio_only: bool = False, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> torch.Tensor: batch_size = audio_hidden_states.size(0) video_batch_size = hidden_states.size(0) @@ -989,7 +994,8 @@ def forward( if self.perturbed_attn: audio_self_attn_args["perturbation_mask"] = perturbation_mask audio_self_attn_args["all_perturbed"] = all_perturbed - attn_audio_hidden_states = self.audio_attn1(**audio_self_attn_args) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_audio_hidden_states = self.audio_attn1(**audio_self_attn_args) audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * audio_gate_msa norm_audio_hidden_states = self.audio_norm2(audio_hidden_states) @@ -997,18 +1003,24 @@ def forward( norm_audio_hidden_states = norm_audio_hidden_states * (1 + audio_scale_text_q) + audio_shift_text_q if self.cross_attn_adaln: audio_encoder_hidden_states = audio_encoder_hidden_states * (1 + audio_scale_text_kv) + audio_shift_text_kv - attn_audio_hidden_states = self.audio_attn2( - norm_audio_hidden_states, - encoder_hidden_states=audio_encoder_hidden_states, - query_rotary_emb=None, - attention_mask=audio_encoder_attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_audio_hidden_states = self.audio_attn2( + norm_audio_hidden_states, + encoder_hidden_states=audio_encoder_hidden_states, + query_rotary_emb=None, + attention_mask=audio_encoder_attention_mask, + ) if self.audio_cross_attn_adaln: attn_audio_hidden_states = attn_audio_hidden_states * audio_gate_text_q audio_hidden_states = audio_hidden_states + attn_audio_hidden_states norm_audio_hidden_states = self.audio_norm3(audio_hidden_states) * (1 + audio_scale_mlp) + audio_shift_mlp - audio_ff_output = self.audio_ff(norm_audio_hidden_states) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + audio_ff_output = checkpoint_fn(self.audio_ff, norm_audio_hidden_states, use_reentrant=False) + else: + audio_ff_output = self.audio_ff(norm_audio_hidden_states) audio_hidden_states = audio_hidden_states + audio_ff_output * audio_gate_mlp return hidden_states, audio_hidden_states @@ -1031,7 +1043,8 @@ def forward( if self.perturbed_attn: video_self_attn_args["perturbation_mask"] = perturbation_mask video_self_attn_args["all_perturbed"] = all_perturbed - attn_hidden_states = self.attn1(**video_self_attn_args) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_hidden_states = self.attn1(**video_self_attn_args) hidden_states = hidden_states + attn_hidden_states * gate_msa norm_audio_hidden_states = self.audio_norm1(audio_hidden_states) @@ -1052,7 +1065,8 @@ def forward( if self.perturbed_attn: audio_self_attn_args["perturbation_mask"] = perturbation_mask audio_self_attn_args["all_perturbed"] = all_perturbed - attn_audio_hidden_states = self.audio_attn1(**audio_self_attn_args) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_audio_hidden_states = self.audio_attn1(**audio_self_attn_args) audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * audio_gate_msa # 2. Video and Audio Cross-Attention with the text embeddings @@ -1067,12 +1081,13 @@ def forward( norm_hidden_states = norm_hidden_states * (1 + scale_text_q) + shift_text_q if self.cross_attn_adaln: encoder_hidden_states = encoder_hidden_states * (1 + scale_text_kv) + shift_text_kv - attn_hidden_states = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - query_rotary_emb=None, - attention_mask=encoder_attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_hidden_states = self.attn2( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states, + query_rotary_emb=None, + attention_mask=encoder_attention_mask, + ) if self.video_cross_attn_adaln: attn_hidden_states = attn_hidden_states * gate_text_q hidden_states = hidden_states + attn_hidden_states @@ -1082,12 +1097,13 @@ def forward( norm_audio_hidden_states = norm_audio_hidden_states * (1 + audio_scale_text_q) + audio_shift_text_q if self.cross_attn_adaln: audio_encoder_hidden_states = audio_encoder_hidden_states * (1 + audio_scale_text_kv) + audio_shift_text_kv - attn_audio_hidden_states = self.audio_attn2( - norm_audio_hidden_states, - encoder_hidden_states=audio_encoder_hidden_states, - query_rotary_emb=None, - attention_mask=audio_encoder_attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_audio_hidden_states = self.audio_attn2( + norm_audio_hidden_states, + encoder_hidden_states=audio_encoder_hidden_states, + query_rotary_emb=None, + attention_mask=audio_encoder_attention_mask, + ) if self.audio_cross_attn_adaln: attn_audio_hidden_states = attn_audio_hidden_states * audio_gate_text_q audio_hidden_states = audio_hidden_states + attn_audio_hidden_states @@ -1119,13 +1135,14 @@ def forward( 1 + audio_a2v_ca_scale.squeeze(2) ) + audio_a2v_ca_shift.squeeze(2) - a2v_attn_hidden_states = self.audio_to_video_attn( - mod_norm_hidden_states, - encoder_hidden_states=mod_norm_audio_hidden_states, - query_rotary_emb=ca_video_rotary_emb, - key_rotary_emb=ca_audio_rotary_emb, - attention_mask=a2v_cross_attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + a2v_attn_hidden_states = self.audio_to_video_attn( + mod_norm_hidden_states, + encoder_hidden_states=mod_norm_audio_hidden_states, + query_rotary_emb=ca_video_rotary_emb, + key_rotary_emb=ca_audio_rotary_emb, + attention_mask=a2v_cross_attention_mask, + ) hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states if use_v2a_cross_attention: @@ -1136,22 +1153,31 @@ def forward( 1 + audio_v2a_ca_scale.squeeze(2) ) + audio_v2a_ca_shift.squeeze(2) - v2a_attn_hidden_states = self.video_to_audio_attn( - mod_norm_audio_hidden_states, - encoder_hidden_states=mod_norm_hidden_states, - query_rotary_emb=ca_audio_rotary_emb, - key_rotary_emb=ca_video_rotary_emb, - attention_mask=v2a_cross_attention_mask, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + v2a_attn_hidden_states = self.video_to_audio_attn( + mod_norm_audio_hidden_states, + encoder_hidden_states=mod_norm_hidden_states, + query_rotary_emb=ca_audio_rotary_emb, + key_rotary_emb=ca_video_rotary_emb, + attention_mask=v2a_cross_attention_mask, + ) audio_hidden_states = audio_hidden_states + v2a_gate * v2a_attn_hidden_states # 4. Feedforward norm_hidden_states = self.norm3(hidden_states) * (1 + scale_mlp) + shift_mlp - ff_output = self.ff(norm_hidden_states) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + ff_output = checkpoint_fn(self.ff, norm_hidden_states, use_reentrant=False) + else: + ff_output = self.ff(norm_hidden_states) hidden_states = hidden_states + ff_output * gate_mlp norm_audio_hidden_states = self.audio_norm3(audio_hidden_states) * (1 + audio_scale_mlp) + audio_shift_mlp - audio_ff_output = self.audio_ff(norm_audio_hidden_states) + if checkpoint_ffn: + audio_ff_output = checkpoint_fn(self.audio_ff, norm_audio_hidden_states, use_reentrant=False) + else: + audio_ff_output = self.audio_ff(norm_audio_hidden_states) audio_hidden_states = audio_hidden_states + audio_ff_output * audio_gate_mlp return hidden_states, audio_hidden_states @@ -1486,6 +1512,8 @@ class LTX2VideoTransformer3DModel( _tread_router: Optional[TREADRouter] = None _tread_routes: Optional[List[Dict[str, Any]]] = None _supports_gradient_checkpointing = True + _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _skip_layerwise_casting_patterns = ["norm"] _repeated_blocks = ["LTX2VideoTransformerBlock"] _cp_plan = { @@ -1731,6 +1759,9 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.time_sign_embed: Optional[nn.Embedding] = None self.audio_time_sign_embed: Optional[nn.Embedding] = None if enable_time_sign_embed: @@ -1768,6 +1799,15 @@ def unfuse_qkv_projections(self): def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: Optional[List[Dict[str, Any]]] = None): """Attach a TREAD router and route definitions.""" self._tread_router = router @@ -2313,8 +2353,81 @@ def _to_pos(idx: int) -> int: if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(self.transformer_blocks, hidden_states.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not musubi_offload_active + and grounding_objs is None + and hidden_states_buffer is None + and not output_hidden_states + ) + segmented_checkpoint_fn = None + segmented_checkpoint_kwargs = {"use_reentrant": False} + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + captured_frame_hidden: Optional[torch.Tensor] = None for block_idx, block in enumerate(self.transformer_blocks): + if use_segmented_checkpointing: + if block_idx != 0: + continue + + def run_segmented_block( + _idx, + block, + hidden_states, + audio_hidden_states, + ): + return block( + hidden_states, + audio_hidden_states, + encoder_hidden_states, + audio_encoder_hidden_states, + temb, + temb_audio, + video_cross_attn_scale_shift, + audio_cross_attn_scale_shift, + video_cross_attn_a2v_gate, + audio_cross_attn_v2a_gate, + temb_prompt, + temb_prompt_audio, + video_rotary_emb, + audio_rotary_emb, + video_cross_attn_rotary_emb, + audio_cross_attn_rotary_emb, + encoder_attention_mask, + audio_encoder_attention_mask, + self_attention_mask, + audio_self_attention_mask, + a2v_cross_attention_mask, + v2a_cross_attention_mask, + True, + True, + None, + None, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + hidden_states, audio_hidden_states = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states, audio_hidden_states), + run_segmented_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + continue + if use_routing and route_ptr < len(routes) and block_idx == routes[route_ptr]["start_layer_idx"]: mask_ratio = routes[route_ptr]["selection_ratio"] tread_mask_info = router.get_mask( @@ -2336,7 +2449,7 @@ def _to_pos(idx: int) -> int: musubi_manager.stream_in(block, hidden_states.device) if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -2388,32 +2501,64 @@ def _checkpoint_block( True, None, None, + offload_attention=self.gradient_checkpointing_offload_attention, ) - checkpoint_kwargs = {"use_reentrant": False} - if checkpoint_fn is simpletuner_checkpoint: - checkpoint_kwargs.update(_transformerengine_checkpoint_kwargs(block)) + if self.gradient_checkpointing_backend.endswith("-ffn"): + hidden_states, audio_hidden_states = block( + hidden_states=hidden_states, + audio_hidden_states=audio_hidden_states, + encoder_hidden_states=encoder_hidden_states, + audio_encoder_hidden_states=audio_encoder_hidden_states, + temb=temb, + temb_audio=temb_audio, + temb_ca_scale_shift=video_cross_attn_scale_shift, + temb_ca_audio_scale_shift=audio_cross_attn_scale_shift, + temb_ca_gate=video_cross_attn_a2v_gate, + temb_ca_audio_gate=audio_cross_attn_v2a_gate, + temb_prompt=temb_prompt, + temb_prompt_audio=temb_prompt_audio, + video_rotary_emb=current_video_rotary_emb, + audio_rotary_emb=audio_rotary_emb, + ca_video_rotary_emb=current_ca_video_rotary_emb, + ca_audio_rotary_emb=audio_cross_attn_rotary_emb, + encoder_attention_mask=encoder_attention_mask, + audio_encoder_attention_mask=audio_encoder_attention_mask, + self_attention_mask=self_attention_mask, + audio_self_attention_mask=audio_self_attention_mask, + a2v_cross_attention_mask=a2v_cross_attention_mask, + v2a_cross_attention_mask=v2a_cross_attention_mask, + use_a2v_cross_attention=True, + use_v2a_cross_attention=True, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + checkpoint_kwargs = {"use_reentrant": False} + if checkpoint_fn is simpletuner_checkpoint: + checkpoint_kwargs.update(_transformerengine_checkpoint_kwargs(block)) - hidden_states, audio_hidden_states = checkpoint_fn( - _checkpoint_block, - hidden_states, - audio_hidden_states, - encoder_hidden_states, - audio_encoder_hidden_states, - temb, - temb_audio, - video_cross_attn_scale_shift, - audio_cross_attn_scale_shift, - video_cross_attn_a2v_gate, - audio_cross_attn_v2a_gate, - temb_prompt, - temb_prompt_audio, - current_video_rotary_emb, - audio_rotary_emb, - current_ca_video_rotary_emb, - audio_cross_attn_rotary_emb, - **checkpoint_kwargs, - ) + hidden_states, audio_hidden_states = checkpoint_fn( + _checkpoint_block, + hidden_states, + audio_hidden_states, + encoder_hidden_states, + audio_encoder_hidden_states, + temb, + temb_audio, + video_cross_attn_scale_shift, + audio_cross_attn_scale_shift, + video_cross_attn_a2v_gate, + audio_cross_attn_v2a_gate, + temb_prompt, + temb_prompt_audio, + current_video_rotary_emb, + audio_rotary_emb, + current_ca_video_rotary_emb, + audio_cross_attn_rotary_emb, + **checkpoint_kwargs, + ) else: hidden_states, audio_hidden_states = block( hidden_states=hidden_states, @@ -2440,6 +2585,7 @@ def _checkpoint_block( v2a_cross_attention_mask=v2a_cross_attention_mask, use_a2v_cross_attention=True, use_v2a_cross_attention=True, + offload_attention=self.gradient_checkpointing_offload_attention, ) if grounding_objs is not None and hasattr(block, "fuser"): diff --git a/simpletuner/helpers/models/lumina2/transformer.py b/simpletuner/helpers/models/lumina2/transformer.py index cb43d1c3a..ce87dcad0 100644 --- a/simpletuner/helpers/models/lumina2/transformer.py +++ b/simpletuner/helpers/models/lumina2/transformer.py @@ -49,6 +49,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.packed_attention_processors import run_packed_qkv_attention logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -605,6 +607,9 @@ def __init__( ) self.gradient_checkpointing = False + self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, @@ -613,6 +618,15 @@ def __init__( logger=logger, ) + def set_gradient_checkpointing_backend(self, backend: str): + self.gradient_checkpointing_backend = backend + + def set_gradient_checkpointing_interval(self, interval: int | None): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def fuse_qkv_projections(self, preferred_backend: Optional[str] = None): for module in self.modules(): if isinstance(module, Attention): @@ -717,15 +731,68 @@ def forward( if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, hidden_states.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not musubi_offload_active + and hidden_states_buffer is None + ) + segmented_checkpoint_fn = None + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + capture_idx = 0 for idx, layer in enumerate(self.layers): + if use_segmented_checkpointing: + if idx != 0: + continue + + def run_segmented_block(_idx, segment_layer, segment_hidden_states): + return segment_layer(segment_hidden_states, attention_mask if use_mask else None, rotary_emb, temb) + + (hidden_states,) = checkpoint_sequential_state( + self.layers, + self.gradient_checkpointing_interval, + (hidden_states,), + run_segmented_block, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + continue + if musubi_offload_active and musubi_manager.is_managed_block(idx): musubi_manager.stream_in(layer, hidden_states.device) - if torch.is_grad_enabled() and self.gradient_checkpointing: - hidden_states = self._gradient_checkpointing_func( - layer, hidden_states, attention_mask if use_mask else None, rotary_emb, temb - ) + if torch.is_grad_enabled() and should_checkpoint_block( + idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + hidden_states = offloaded_checkpoint( + layer, + hidden_states, + attention_mask if use_mask else None, + rotary_emb, + temb, + use_reentrant=False, + ) + else: + hidden_states = self._gradient_checkpointing_func( + layer, hidden_states, attention_mask if use_mask else None, rotary_emb, temb + ) else: hidden_states = layer(hidden_states, attention_mask if use_mask else None, rotary_emb, temb) diff --git a/simpletuner/helpers/models/mageflow/model.py b/simpletuner/helpers/models/mageflow/model.py index c8c99ed72..52f8fe20f 100644 --- a/simpletuner/helpers/models/mageflow/model.py +++ b/simpletuner/helpers/models/mageflow/model.py @@ -2,6 +2,7 @@ from typing import Optional import torch +from diffusers import FlowMatchEulerDiscreteScheduler from einops import rearrange from transformers import AutoProcessor, AutoTokenizer, Qwen3VLForConditionalGeneration @@ -22,6 +23,7 @@ from simpletuner.helpers.models.mageflow.vendor.pipeline import _lens_to_cu from simpletuner.helpers.models.registry import ModelRegistry from simpletuner.helpers.musubi_block_swap import apply_musubi_pretrained_defaults +from simpletuner.helpers.training.flow_match import fix_flow_match_euler_schedule_bounds logger = logging.getLogger(__name__) @@ -49,12 +51,12 @@ class MageFlow(ImageModelFoundation): DEFAULT_MODEL_FLAVOUR = "base" HUGGINGFACE_PATHS = { - "base": "microsoft/Mage-Flow-Base", - "default": "microsoft/Mage-Flow", - "turbo": "microsoft/Mage-Flow-Turbo", - "edit-base": "microsoft/Mage-Flow-Edit-Base", - "edit": "microsoft/Mage-Flow-Edit", - "edit-turbo": "microsoft/Mage-Flow-Edit-Turbo", + "base": "natalie5/Mage-Flow-Base", + "default": "natalie5/Mage-Flow", + "turbo": "natalie5/Mage-Flow-Turbo", + "edit-base": "natalie5/Mage-Flow-Edit-Base", + "edit": "natalie5/Mage-Flow-Edit", + "edit-turbo": "natalie5/Mage-Flow-Edit-Turbo", } MODEL_LICENSE = "mit" @@ -188,6 +190,52 @@ def _model_flavour(self) -> str: def _is_edit_flavour(self) -> bool: return self._model_flavour() in {"edit-base", "edit", "edit-turbo"} + def _resolve_checkpointing_transformer(self): + seen = set() + + def visit(module): + if module is None or id(module) in seen: + return None + seen.add(id(module)) + try: + module = self.unwrap_model(model=module) + except Exception: + pass + if isinstance(module, MageFlowTransformer2DModel): + return module + for attr_name in ("base_model", "model", "module", "_orig_mod"): + child = getattr(module, attr_name, None) + if child is not module: + resolved = visit(child) + if resolved is not None: + return resolved + return None + + return visit(getattr(self, "model", None)) + + def enable_gradient_checkpointing(self): + transformer = self._resolve_checkpointing_transformer() + if transformer is not None: + transformer.enable_gradient_checkpointing() + + def disable_gradient_checkpointing(self): + transformer = self._resolve_checkpointing_transformer() + if transformer is not None: + transformer.disable_gradient_checkpointing() + + def setup_training_noise_schedule(self): + try: + super().setup_training_noise_schedule() + except OSError: + logger.info("Mage-Flow checkpoint has no scheduler config; using the default training flow scheduler.") + self.noise_schedule = FlowMatchEulerDiscreteScheduler( + num_train_timesteps=1000, + shift=self.config.flow_schedule_shift if self.config.flow_schedule_shift is not None else 6.0, + use_dynamic_shifting=False, + ) + fix_flow_match_euler_schedule_bounds(self.noise_schedule) + return self.config, self.noise_schedule + @staticmethod def _normalise_attention_mechanism(attention_mechanism: str | None) -> str: return str(attention_mechanism or "").strip().lower().replace("_", "-") @@ -338,8 +386,8 @@ def _select_crepa_hidden_states(self, prepared_batch: dict, hidden_states_buffer def _load_processor_for_pipeline(self): if self.processor is not None: return self.processor - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) self.processor = self.PROCESSOR_CLASS.from_pretrained( processor_path, subfolder=processor_subfolder, diff --git a/simpletuner/helpers/models/mageflow/pipeline.py b/simpletuner/helpers/models/mageflow/pipeline.py index 406ead58e..6e575bb14 100644 --- a/simpletuner/helpers/models/mageflow/pipeline.py +++ b/simpletuner/helpers/models/mageflow/pipeline.py @@ -24,6 +24,14 @@ def _resolve_repo_dir(repo_id_or_path: str, *, revision: Optional[str] = None, l return snapshot_download(repo_id=repo_id_or_path, revision=revision, local_files_only=local_files_only) +def _load_scheduler(repo_dir: str): + scheduler_dir = os.path.join(repo_dir, "scheduler") + scheduler_config = os.path.join(scheduler_dir, "scheduler_config.json") + if os.path.isdir(scheduler_dir) and os.path.exists(scheduler_config): + return FlowMatchEulerDiscreteScheduler.from_pretrained(scheduler_dir) + return FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000, shift=6.0, use_dynamic_shifting=False) + + def _call_vae_decode(vae, latents: torch.Tensor) -> torch.Tensor: if hasattr(vae, "decode_to_tensor"): return vae.decode_to_tensor(latents) @@ -236,7 +244,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: str, **kwargs): revision=revision, local_files_only=local_files_only, ) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(os.path.join(repo_dir, "scheduler")) + scheduler = _load_scheduler(repo_dir) return cls( transformer=transformer, vae=vae, diff --git a/simpletuner/helpers/models/mageflow/transformer.py b/simpletuner/helpers/models/mageflow/transformer.py index 664a2404c..337a438c5 100644 --- a/simpletuner/helpers/models/mageflow/transformer.py +++ b/simpletuner/helpers/models/mageflow/transformer.py @@ -14,6 +14,7 @@ from simpletuner.helpers.models.mageflow.vendor.models.modules._attn_backend import set_attn_backend from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager from simpletuner.helpers.training.gradient_checkpointing_interval import ( + checkpoint_sequential_state, get_checkpoint_backend, get_checkpoint_backend_scope, get_checkpoint_function, @@ -51,6 +52,7 @@ def _resolve_repo_dir(repo_id_or_path: str, *, revision: Optional[str] = None, l class MageFlowTransformer2DModel(MageFlow, ModelMixin, ConfigMixin, PeftAdapterMixin): _supports_gradient_checkpointing = True _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _no_split_modules = ["MageFlowTransformerBlock"] _skip_layerwise_casting_patterns = ["pos_embed", "norm"] @@ -128,6 +130,9 @@ def __init__( torch.nn.init.zeros_(self.time_sign_embed.weight) self.gradient_checkpointing_backend = get_checkpoint_backend() self.gradient_checkpointing_scope = get_checkpoint_backend_scope() + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None + self.gradient_checkpointing_offload_attention = False self._gradient_checkpointing_func = get_checkpoint_function() self._musubi_block_swap = MusubiBlockSwapManager.build( depth=depth, @@ -176,6 +181,15 @@ def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_scope = get_checkpoint_backend_scope(backend) self._gradient_checkpointing_func = get_checkpoint_function() + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def enable_gradient_checkpointing(self, gradient_checkpointing_func=None): self.checkpoint = True self.config.checkpoint = True @@ -256,6 +270,45 @@ def forward( skip_layers_set = set(skip_layers) if skip_layers is not None else set() for index_block, block in enumerate(self.transformer_blocks): + if ( + self.training + and self.checkpoint + and self.gradient_checkpointing_scope == "layer" + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not musubi_offload_active + and not skip_layers_set + and hidden_states_buffer is None + ): + if index_block != 0: + continue + for segment_block in self.transformer_blocks: + _ensure_module_device(getattr(segment_block, "img_mod", None), img.device) + _ensure_module_device(getattr(segment_block, "txt_mod", None), img.device) + + def run_segment_block(_relative_index, segment_block, segment_txt, segment_img): + return segment_block( + hidden_states=segment_img, + encoder_hidden_states=segment_txt, + txt_cu_lens=txt_cu_seqlens, + img_cu_lens=img_cu_seqlens, + temb=temb, + image_rotary_emb=ms_pe, + joint_attention_kwargs=attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + txt, img = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (txt, img), + run_segment_block, + self._gradient_checkpointing_func, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + continue + _ensure_module_device(getattr(block, "img_mod", None), img.device) _ensure_module_device(getattr(block, "txt_mod", None), img.device) if musubi_offload_active and musubi_manager.is_managed_block(index_block): @@ -273,10 +326,32 @@ def forward( joint_attention_kwargs=attention_kwargs, checkpoint_ffn=True, checkpoint_fn=self._gradient_checkpointing_func, + offload_attention=self.gradient_checkpointing_offload_attention, ) elif self.training and self.checkpoint: + + def run_checkpointed_block( + checkpoint_img, + checkpoint_txt, + checkpoint_temb, + checkpoint_pe, + checkpoint_txt_cu_seqlens, + checkpoint_img_cu_seqlens, + checkpoint_block=block, + ): + return checkpoint_block( + hidden_states=checkpoint_img, + encoder_hidden_states=checkpoint_txt, + txt_cu_lens=checkpoint_txt_cu_seqlens, + img_cu_lens=checkpoint_img_cu_seqlens, + temb=checkpoint_temb, + image_rotary_emb=checkpoint_pe, + joint_attention_kwargs=attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + txt, img = self._gradient_checkpointing_func( - block, + run_checkpointed_block, img, txt, temb, @@ -294,6 +369,7 @@ def forward( temb=temb, image_rotary_emb=ms_pe, joint_attention_kwargs=attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) if musubi_offload_active and musubi_manager.is_managed_block(index_block): musubi_manager.stream_out(block) diff --git a/simpletuner/helpers/models/mageflow/vendor/models/modules/mage_layers.py b/simpletuner/helpers/models/mageflow/vendor/models/modules/mage_layers.py index d79188e1f..b995d2a60 100644 --- a/simpletuner/helpers/models/mageflow/vendor/models/modules/mage_layers.py +++ b/simpletuner/helpers/models/mageflow/vendor/models/modules/mage_layers.py @@ -10,6 +10,8 @@ from torch import Tensor from torch._dynamo import allow_in_graph as maybe_allow_in_graph +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context + from ._attn_backend import flash_attn_varlen_func @@ -619,6 +621,7 @@ def forward( joint_attention_kwargs: dict[str, Any] | None = None, checkpoint_ffn: bool = False, checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: # Get modulation parameters for both streams # if isinstance(temb, tuple): @@ -658,17 +661,18 @@ def forward( joint_attention_kwargs = joint_attention_kwargs or {} # logger.info(f"img_modulated: {img_modulated}") # logger.info(f"txt_modulated: {txt_modulated}") - attn_output = self.attn( - hidden_states=img_modulated, # Image stream (will be processed as "sample") - encoder_hidden_states=txt_modulated, # Text stream (will be processed as "context") - # encoder_hidden_states_mask=encoder_hidden_states_mask, - image_rotary_emb=image_rotary_emb, - txt_cu_lens=txt_cu_lens, - img_cu_lens=img_cu_lens, - # freqs_cos=freqs_cos, - # freqs_sin=freqs_sin, - **joint_attention_kwargs, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output = self.attn( + hidden_states=img_modulated, # Image stream (will be processed as "sample") + encoder_hidden_states=txt_modulated, # Text stream (will be processed as "context") + # encoder_hidden_states_mask=encoder_hidden_states_mask, + image_rotary_emb=image_rotary_emb, + txt_cu_lens=txt_cu_lens, + img_cu_lens=img_cu_lens, + # freqs_cos=freqs_cos, + # freqs_sin=freqs_sin, + **joint_attention_kwargs, + ) # logger.info(f"attn_output: {attn_output}") # MageDoubleStreamAttnProcessor returns (img_output, txt_output) when encoder_hidden_states is provided diff --git a/simpletuner/helpers/models/mageflow/vendor/pipeline.py b/simpletuner/helpers/models/mageflow/vendor/pipeline.py index d2a8787f2..cf5f6e5d7 100644 --- a/simpletuner/helpers/models/mageflow/vendor/pipeline.py +++ b/simpletuner/helpers/models/mageflow/vendor/pipeline.py @@ -910,5 +910,14 @@ def _resolve(p): model.vae.to(torch.bfloat16) model.eval() # Diffusers FlowMatchEulerDiscreteScheduler (scheduler/scheduler_config.json). - model.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(_safe_subpath(repo_dir, "scheduler")) + scheduler_dir = _safe_subpath(repo_dir, "scheduler") + scheduler_config = os.path.join(scheduler_dir, "scheduler_config.json") + if os.path.isdir(scheduler_dir) and os.path.exists(scheduler_config): + model.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(scheduler_dir) + else: + model.scheduler = build_scheduler( + 1, + device=device, + shift=(cfg.static_shift if cfg.static_shift is not None else 6.0), + ) return model diff --git a/simpletuner/helpers/models/omnigen/model.py b/simpletuner/helpers/models/omnigen/model.py index cd23c34d5..e7208164f 100644 --- a/simpletuner/helpers/models/omnigen/model.py +++ b/simpletuner/helpers/models/omnigen/model.py @@ -161,6 +161,9 @@ def __init__(self, config, accelerator): def supports_crepa_self_flow(self) -> bool: return True + def uses_text_embeddings_cache(self) -> bool: + return False + def get_transforms(self, dataset_type: str = "image"): from torchvision import transforms diff --git a/simpletuner/helpers/models/pixart/transformer.py b/simpletuner/helpers/models/pixart/transformer.py index 6f923ebaa..bcc1ca30b 100644 --- a/simpletuner/helpers/models/pixart/transformer.py +++ b/simpletuner/helpers/models/pixart/transformer.py @@ -36,6 +36,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.packed_attention_processors import PackedFusedAttnProcessor2_0 from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import CallableDict, MutableModuleList, PatchableModule @@ -350,6 +352,8 @@ def __init__( self.caption_projection = PixArtAlphaTextProjection( in_features=self.config.caption_channels, hidden_size=self.inner_dim ) + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, @@ -361,6 +365,12 @@ def __init__( def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + @property # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors def attn_processors(self) -> Dict[str, AttentionProcessor]: @@ -614,8 +624,54 @@ def _to_pos(idx): if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, hidden_states.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not musubi_offload_active + and controlnet_block_samples is None + and hidden_states_buffer is None + ) + segmented_checkpoint_fn = None + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + capture_idx = 0 for index_block, block in enumerate(self.transformer_blocks): + if use_segmented_checkpointing: + if index_block != 0: + continue + + def run_segmented_block(_idx, segment_block, segment_hidden_states): + return segment_block( + segment_hidden_states, + attention_mask=attention_mask, + encoder_hidden_states=encoder_hidden_states, + encoder_attention_mask=encoder_attention_mask, + timestep=timestep, + cross_attention_kwargs=cross_attention_kwargs, + class_labels=None, + ) + + (hidden_states,) = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_segmented_block, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + continue + if musubi_offload_active and musubi_manager.is_managed_block(index_block): musubi_manager.stream_in(block, hidden_states.device) # TREAD: START a route? @@ -630,8 +686,13 @@ def _to_pos(idx): hidden_states = router.start_route(hidden_states, tread_mask_info) routing_now = True - if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if torch.is_grad_enabled() and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/qwen_image/model.py b/simpletuner/helpers/models/qwen_image/model.py index 9b4a09179..f36e2f0eb 100644 --- a/simpletuner/helpers/models/qwen_image/model.py +++ b/simpletuner/helpers/models/qwen_image/model.py @@ -245,8 +245,8 @@ def _load_processor_for_pipeline(self): if processor_cls is None: return None - processor_path = getattr(self.config, "processor_pretrained_model_name_or_path", None) or self._model_config_path() - processor_subfolder = getattr(self.config, "processor_subfolder", self.PROCESSOR_SUBFOLDER) + processor_path = self._resolve_qwen_processor_path(self._model_config_path()) + processor_subfolder = self._resolve_qwen_processor_subfolder(self.PROCESSOR_SUBFOLDER) processor_revision = getattr(self.config, "processor_revision", getattr(self.config, "revision", None)) processor_kwargs = {"pretrained_model_name_or_path": processor_path} diff --git a/simpletuner/helpers/models/qwen_image/transformer.py b/simpletuner/helpers/models/qwen_image/transformer.py index 25c1acd2d..bbfe63ada 100644 --- a/simpletuner/helpers/models/qwen_image/transformer.py +++ b/simpletuner/helpers/models/qwen_image/transformer.py @@ -43,6 +43,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import MutableModuleList, PatchableModule @@ -1099,6 +1100,8 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, blocks_to_swap=musubi_blocks_to_swap, @@ -1113,6 +1116,12 @@ def __init__( def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: Optional[List[Dict]] = None): """Set TREAD router and routes for token reduction during training.""" self._tread_router = router @@ -1373,7 +1382,12 @@ def _to_pos(idx: int) -> int: hidden_states = router.start_route(hidden_states, tread_mask_info) routing_now = True - if torch.is_grad_enabled() and self.gradient_checkpointing: + if torch.is_grad_enabled() and should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): def create_custom_forward(module, mod_idx): def custom_forward( @@ -1407,7 +1421,7 @@ def custom_forward( return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/sana/transformer.py b/simpletuner/helpers/models/sana/transformer.py index cc93c61ae..b21370ba7 100644 --- a/simpletuner/helpers/models/sana/transformer.py +++ b/simpletuner/helpers/models/sana/transformer.py @@ -35,6 +35,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import CallableDict, MutableModuleList, PatchableModule @@ -472,6 +473,7 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" # tread support @@ -500,6 +502,9 @@ def set_gradient_checkpointing_interval(self, interval: int): """ self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend @@ -742,7 +747,12 @@ def forward( if ( self.training and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or i % self.gradient_checkpointing_interval == 0) + and should_checkpoint_block( + i, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) ): def create_custom_forward(module): @@ -751,7 +761,7 @@ def custom_forward(*inputs): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/sanavideo/transformer.py b/simpletuner/helpers/models/sanavideo/transformer.py index a3267ce5a..40696e3e4 100644 --- a/simpletuner/helpers/models/sanavideo/transformer.py +++ b/simpletuner/helpers/models/sanavideo/transformer.py @@ -39,6 +39,8 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits @@ -769,6 +771,9 @@ def __init__( self.proj_out = nn.Linear(inner_dim, math.prod(patch_size) * out_channels) self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None + self.gradient_checkpointing_backend = "torch" self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, @@ -777,6 +782,15 @@ def __init__( logger=logger, ) + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + + def set_gradient_checkpointing_backend(self, backend: str): + self.gradient_checkpointing_backend = backend + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: if hasattr(self.time_embed, "enable_flowmap_time_conditioning"): self.time_embed.enable_flowmap_time_conditioning(gate_value=gate_value, deltatime_type=deltatime_type) @@ -988,13 +1002,29 @@ def forward( musubi_offload_active = musubi_manager.activate(combined_blocks, hidden_states.device, grad_enabled) capture_idx = 0 - if torch.is_grad_enabled() and self.gradient_checkpointing: - for index_block, block in enumerate(self.transformer_blocks): - if musubi_offload_active and musubi_manager.is_managed_block(index_block): - musubi_manager.stream_in(block, hidden_states.device) - hidden_states = self._gradient_checkpointing_func( - block, - hidden_states, + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not musubi_offload_active + and grounding_objs is None + and controlnet_block_samples is None + and not output_hidden_states + and hidden_states_buffer is None + ) + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + + def run_segmented_block(_idx, segment_block, segment_hidden_states): + return segment_block( + segment_hidden_states, attention_mask, encoder_hidden_states, encoder_attention_mask, @@ -1004,6 +1034,57 @@ def forward( post_patch_width, rotary_emb, ) + + (hidden_states,) = checkpoint_sequential_state( + self.transformer_blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_segmented_block, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + elif torch.is_grad_enabled() and self.gradient_checkpointing: + for index_block, block in enumerate(self.transformer_blocks): + if musubi_offload_active and musubi_manager.is_managed_block(index_block): + musubi_manager.stream_in(block, hidden_states.device) + if should_checkpoint_block( + index_block, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + checkpoint_fn = offloaded_checkpoint + else: + checkpoint_fn = self._gradient_checkpointing_func + + hidden_states = checkpoint_fn( + block, + hidden_states, + attention_mask, + encoder_hidden_states, + encoder_attention_mask, + timestep, + post_patch_num_frames, + post_patch_height, + post_patch_width, + rotary_emb, + ) + else: + hidden_states = block( + hidden_states, + attention_mask, + encoder_hidden_states, + encoder_attention_mask, + timestep, + post_patch_num_frames, + post_patch_height, + post_patch_width, + rotary_emb, + ) if grounding_objs is not None and hasattr(block, "fuser"): hidden_states = apply_grounding_fuser( block.fuser, diff --git a/simpletuner/helpers/models/sd3/expanded.py b/simpletuner/helpers/models/sd3/expanded.py index cabb4e12b..1292123dd 100644 --- a/simpletuner/helpers/models/sd3/expanded.py +++ b/simpletuner/helpers/models/sd3/expanded.py @@ -404,9 +404,18 @@ def unfuse_qkv_projections(self): if self.original_attn_processors is not None: self.set_attn_processor(self.original_attn_processors) - def _set_gradient_checkpointing(self, module, value=False): + def _set_gradient_checkpointing(self, module=None, value=False, enable=None, gradient_checkpointing_func=None): + if enable is not None: + value = enable + if gradient_checkpointing_func is not None: + self._gradient_checkpointing_func = gradient_checkpointing_func + if module is None: + module = self + self.gradient_checkpointing = value if hasattr(module, "gradient_checkpointing"): module.gradient_checkpointing = value + for child in module.children(): + self._set_gradient_checkpointing(child, value=value) def forward( self, @@ -474,7 +483,7 @@ def custom_forward(*inputs): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/models/sd3/transformer.py b/simpletuner/helpers/models/sd3/transformer.py index ff86260ae..c7daa9c0c 100644 --- a/simpletuner/helpers/models/sd3/transformer.py +++ b/simpletuner/helpers/models/sd3/transformer.py @@ -38,7 +38,10 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.checkpointing import checkpoint as simpletuner_checkpoint +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.packed_attention_processors import PackedJointAttnProcessor2_0 from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import CallableDict, MutableModuleList, PatchableModule @@ -146,16 +149,9 @@ def _sd3_apply_joint_transformer_block( temb_hidden: torch.Tensor, temb_context: torch.Tensor, joint_attention_kwargs: Optional[Dict[str, Any]] = None, + offload_attention: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: joint_attention_kwargs = joint_attention_kwargs or {} - if temb_hidden.ndim == 2 and temb_context.ndim == 2: - return block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb_hidden, - joint_attention_kwargs=joint_attention_kwargs, - ) - if block.use_dual_attention: ( norm_hidden_states, @@ -182,11 +178,12 @@ def _sd3_apply_joint_transformer_block( c_gate_mlp, ) = _sd3_apply_ada_layer_norm_zero(block.norm1_context, encoder_hidden_states, temb_context) - attn_output, context_attn_output = block.attn( - hidden_states=norm_hidden_states, - encoder_hidden_states=norm_encoder_hidden_states, - **joint_attention_kwargs, - ) + with activation_offload_context(offload_attention, label=f"{block.__class__.__qualname__}:attention"): + attn_output, context_attn_output = block.attn( + hidden_states=norm_hidden_states, + encoder_hidden_states=norm_encoder_hidden_states, + **joint_attention_kwargs, + ) if gate_msa.ndim == 2: gate_msa = gate_msa.unsqueeze(1) @@ -194,7 +191,8 @@ def _sd3_apply_joint_transformer_block( hidden_states = hidden_states + attn_output if block.use_dual_attention: - attn_output2 = block.attn2(hidden_states=norm_hidden_states2, **joint_attention_kwargs) + with activation_offload_context(offload_attention, label=f"{block.__class__.__qualname__}:dual_attention"): + attn_output2 = block.attn2(hidden_states=norm_hidden_states2, **joint_attention_kwargs) if gate_msa2.ndim == 2: gate_msa2 = gate_msa2.unsqueeze(1) attn_output2 = gate_msa2 * attn_output2 @@ -269,6 +267,7 @@ class SD3Transformer2DModel(PatchableModule, ModelMixin, ConfigMixin, PeftAdapte "PatchEmbed", ] _supports_gradient_checkpointing = True + _supports_attention_activation_offload = True _fsdp_exclude_auto_wrap_modules = ["PatchEmbed"] _cp_plan = { "": { @@ -375,7 +374,9 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, @@ -387,9 +388,15 @@ def __init__( def set_gradient_checkpointing_interval(self, interval: int): self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): self._tread_router = router self._tread_routes = routes @@ -719,116 +726,152 @@ def _to_pos(idx): if hasattr(self, "position_net") and grounding_kwargs is not None: grounding_objs = self.position_net(**grounding_kwargs) - capture_idx = 0 - for index_block, block in enumerate(self.transformer_blocks): - if musubi_offload_active and musubi_manager.is_managed_block(index_block): - musubi_manager.stream_in(block, hidden_states.device) - # TREAD: START a route? - if use_routing and route_ptr < len(routes) and global_idx == routes[route_ptr]["start_layer_idx"]: - mask_ratio = routes[route_ptr]["selection_ratio"] - tread_mask_info = router.get_mask( - hidden_states, - mask_ratio=mask_ratio, - force_keep=force_keep_mask, - ) - saved_tokens = hidden_states.clone() - hidden_states = router.start_route(hidden_states, tread_mask_info) - routing_now = True - - # Skip specified layers - if skip_layers is not None and index_block in skip_layers: - if block_controlnet_hidden_states is not None and block.context_pre_only is False: - interval_control = len(self.transformer_blocks) // len(block_controlnet_hidden_states) - hidden_states = hidden_states + block_controlnet_hidden_states[index_block // interval_control] - continue + segment_size = self.gradient_checkpointing_interval + use_segmented_checkpointing = ( + self.training + and self.gradient_checkpointing + and segment_size is not None + and segment_size > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not musubi_offload_active + and skip_layers is None + and block_controlnet_hidden_states is None + and grounding_objs is None + and hidden_states_buffer is None + ) + segmented_checkpoint_fn = None + segmented_checkpoint_kwargs: Dict[str, Any] = {} + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint - if ( - self.training - and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or index_block % self.gradient_checkpointing_interval == 0) + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = simpletuner_checkpoint + segmented_checkpoint_kwargs = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} + + def run_segment_block( + _relative_index, + segment_block, + segment_encoder_hidden_states, + segment_hidden_states, ): + return _sd3_apply_joint_transformer_block( + segment_block, + segment_hidden_states, + segment_encoder_hidden_states, + temb_hidden, + temb_context, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) - def create_custom_forward(module, return_dict=None): - def custom_forward(*inputs): - if return_dict is not None: - return module(*inputs, return_dict=return_dict) - else: - return module(*inputs) - - return custom_forward + encoder_hidden_states, hidden_states = checkpoint_sequential_state( + self.transformer_blocks, + segment_size, + (encoder_hidden_states, hidden_states), + run_segment_block, + segmented_checkpoint_fn, + segmented_checkpoint_kwargs, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + else: + capture_idx = 0 + for index_block, block in enumerate(self.transformer_blocks): + if musubi_offload_active and musubi_manager.is_managed_block(index_block): + musubi_manager.stream_in(block, hidden_states.device) + # TREAD: START a route? + if use_routing and route_ptr < len(routes) and global_idx == routes[route_ptr]["start_layer_idx"]: + mask_ratio = routes[route_ptr]["selection_ratio"] + tread_mask_info = router.get_mask( + hidden_states, + mask_ratio=mask_ratio, + force_keep=force_keep_mask, + ) + saved_tokens = hidden_states.clone() + hidden_states = router.start_route(hidden_states, tread_mask_info) + routing_now = True + + # Skip specified layers + if skip_layers is not None and index_block in skip_layers: + if block_controlnet_hidden_states is not None and block.context_pre_only is False: + interval_control = len(self.transformer_blocks) // len(block_controlnet_hidden_states) + hidden_states = hidden_states + block_controlnet_hidden_states[index_block // interval_control] + continue + + if ( + self.training + and self.gradient_checkpointing + and should_checkpoint_block( + index_block, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + ): + + def custom_forward(*inputs, checkpoint_block=block, attention_kwargs=joint_attention_kwargs): + return _sd3_apply_joint_transformer_block( + checkpoint_block, + inputs[0], + inputs[1], + inputs[2], + inputs[3], + joint_attention_kwargs=attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) - if self.gradient_checkpointing_backend == "unsloth": - from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint - checkpoint_fn = offloaded_checkpoint - else: - checkpoint_fn = torch.utils.checkpoint.checkpoint + checkpoint_fn = offloaded_checkpoint + else: + checkpoint_fn = simpletuner_checkpoint - ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} - if temb_hidden.ndim == 2: + ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {} encoder_hidden_states, hidden_states = checkpoint_fn( - create_custom_forward(block), + custom_forward, hidden_states, encoder_hidden_states, temb_hidden, + temb_context, **ckpt_kwargs, ) else: - - def create_tokenwise_forward(module, attention_kwargs): - def custom_forward(*inputs): - return _sd3_apply_joint_transformer_block( - module, - inputs[0], - inputs[1], - inputs[2], - inputs[3], - joint_attention_kwargs=attention_kwargs, - ) - - return custom_forward - - encoder_hidden_states, hidden_states = checkpoint_fn( - create_tokenwise_forward(block, joint_attention_kwargs), + encoder_hidden_states, hidden_states = _sd3_apply_joint_transformer_block( + block, hidden_states, encoder_hidden_states, temb_hidden, temb_context, - **ckpt_kwargs, + joint_attention_kwargs=joint_attention_kwargs, + offload_attention=self.gradient_checkpointing_offload_attention, ) - else: - encoder_hidden_states, hidden_states = _sd3_apply_joint_transformer_block( - block, - hidden_states, - encoder_hidden_states, - temb_hidden, - temb_context, - joint_attention_kwargs=joint_attention_kwargs, - ) - if grounding_objs is not None and hasattr(block, "fuser"): - hidden_states = apply_grounding_fuser(block.fuser, hidden_states, grounding_objs) + if grounding_objs is not None and hasattr(block, "fuser"): + hidden_states = apply_grounding_fuser(block.fuser, hidden_states, grounding_objs) - # controlnet residual - if block_controlnet_hidden_states is not None and block.context_pre_only is False: - interval_control = len(self.transformer_blocks) // len(block_controlnet_hidden_states) - hidden_states = hidden_states + block_controlnet_hidden_states[index_block // interval_control] + # controlnet residual + if block_controlnet_hidden_states is not None and block.context_pre_only is False: + interval_control = len(self.transformer_blocks) // len(block_controlnet_hidden_states) + hidden_states = hidden_states + block_controlnet_hidden_states[index_block // interval_control] - # TREAD: END the current route? - if routing_now and global_idx == routes[route_ptr]["end_layer_idx"]: - hidden_states = router.end_route( - hidden_states, - tread_mask_info, - original_x=saved_tokens, - ) - routing_now = False - route_ptr += 1 - - if musubi_offload_active and musubi_manager.is_managed_block(index_block): - musubi_manager.stream_out(block) - _store_hidden_state(hidden_states_buffer, f"layer_{capture_idx}", hidden_states) - capture_idx += 1 - global_idx += 1 + # TREAD: END the current route? + if routing_now and global_idx == routes[route_ptr]["end_layer_idx"]: + hidden_states = router.end_route( + hidden_states, + tread_mask_info, + original_x=saved_tokens, + ) + routing_now = False + route_ptr += 1 + + if musubi_offload_active and musubi_manager.is_managed_block(index_block): + musubi_manager.stream_out(block) + _store_hidden_state(hidden_states_buffer, f"layer_{capture_idx}", hidden_states) + capture_idx += 1 + global_idx += 1 hidden_states = _sd3_apply_ada_layer_norm_continuous(self.norm_out, hidden_states, temb_hidden) hidden_states = self.proj_out(hidden_states) diff --git a/simpletuner/helpers/models/stable_cascade/unet.py b/simpletuner/helpers/models/stable_cascade/unet.py index ab6d0fefb..c6e6a54de 100644 --- a/simpletuner/helpers/models/stable_cascade/unet.py +++ b/simpletuner/helpers/models/stable_cascade/unet.py @@ -32,6 +32,7 @@ register_flowmap_config, validate_flowmap_deltatime_type, ) +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block # Copied from diffusers.pipelines.wuerstchen.modeling_wuerstchen_common.WuerstchenLayerNorm with WuerstchenLayerNorm -> SDCascadeLayerNorm @@ -446,6 +447,8 @@ def get_block(block_type, in_channels, nhead, c_skip=0, dropout=0, self_attn=Tru ) self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.flowmap_deltatime_type: Optional[str] = None if deltatime_type is not None: self.enable_flowmap_time_conditioning( @@ -453,6 +456,12 @@ def get_block(block_type, in_channels, nhead, c_skip=0, dropout=0, self_attn=Tru deltatime_type=deltatime_type, ) + def set_gradient_checkpointing_interval(self, interval: int | None) -> None: + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None) -> None: + self.gradient_checkpointing_segment_stride = segment_stride + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.flowmap_deltatime_type = validate_flowmap_deltatime_type(deltatime_type, model_name="Stable Cascade") for module in self.modules(): @@ -519,92 +528,69 @@ def get_clip_embeddings(self, clip_txt_pooled, clip_txt=None, clip_img=None): clip = clip_txt_pool return self.clip_norm(clip) + def _run_checkpointable_block(self, block_index: int, block: nn.Module, *args): + if ( + torch.is_grad_enabled() + and self.gradient_checkpointing + and should_checkpoint_block( + block_index, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) + ): + return self._gradient_checkpointing_func(block, *args) + return block(*args) + def _down_encode(self, x, r_embed, clip, delta_r_embed=None): level_outputs = [] block_group = zip(self.down_blocks, self.down_downscalers, self.down_repeat_mappers) - - if torch.is_grad_enabled() and self.gradient_checkpointing: - for down_block, downscaler, repmap in block_group: - x = downscaler(x) - for i in range(len(repmap) + 1): - for block in down_block: - if isinstance(block, SDCascadeResBlock): - x = self._gradient_checkpointing_func(block, x) - elif isinstance(block, SDCascadeAttnBlock): - x = self._gradient_checkpointing_func(block, x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = self._gradient_checkpointing_func(block, x, r_embed, delta_r_embed) - else: - x = self._gradient_checkpointing_func(block) - if i < len(repmap): - x = repmap[i](x) - level_outputs.insert(0, x) - else: - for down_block, downscaler, repmap in block_group: - x = downscaler(x) - for i in range(len(repmap) + 1): - for block in down_block: - if isinstance(block, SDCascadeResBlock): - x = block(x) - elif isinstance(block, SDCascadeAttnBlock): - x = block(x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = block(x, r_embed, delta_r_embed) - else: - x = block(x) - if i < len(repmap): - x = repmap[i](x) - level_outputs.insert(0, x) - return level_outputs - - def _up_decode(self, level_outputs, r_embed, clip, delta_r_embed=None): + checkpoint_block_index = 0 + + for down_block, downscaler, repmap in block_group: + x = downscaler(x) + for i in range(len(repmap) + 1): + for block in down_block: + if isinstance(block, SDCascadeResBlock): + x = self._run_checkpointable_block(checkpoint_block_index, block, x) + elif isinstance(block, SDCascadeAttnBlock): + x = self._run_checkpointable_block(checkpoint_block_index, block, x, clip) + elif isinstance(block, SDCascadeTimestepBlock): + x = self._run_checkpointable_block(checkpoint_block_index, block, x, r_embed, delta_r_embed) + else: + x = self._run_checkpointable_block(checkpoint_block_index, block, x) + checkpoint_block_index += 1 + if i < len(repmap): + x = repmap[i](x) + level_outputs.insert(0, x) + return level_outputs, checkpoint_block_index + + def _up_decode(self, level_outputs, r_embed, clip, delta_r_embed=None, checkpoint_block_index: int = 0): x = level_outputs[0] block_group = zip(self.up_blocks, self.up_upscalers, self.up_repeat_mappers) - if torch.is_grad_enabled() and self.gradient_checkpointing: - for i, (up_block, upscaler, repmap) in enumerate(block_group): - for j in range(len(repmap) + 1): - for k, block in enumerate(up_block): - if isinstance(block, SDCascadeResBlock): - skip = level_outputs[i] if k == 0 and i > 0 else None - if skip is not None and (x.size(-1) != skip.size(-1) or x.size(-2) != skip.size(-2)): - orig_type = x.dtype - x = torch.nn.functional.interpolate( - x.float(), skip.shape[-2:], mode="bilinear", align_corners=True - ) - x = x.to(orig_type) - x = self._gradient_checkpointing_func(block, x, skip) - elif isinstance(block, SDCascadeAttnBlock): - x = self._gradient_checkpointing_func(block, x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = self._gradient_checkpointing_func(block, x, r_embed, delta_r_embed) - else: - x = self._gradient_checkpointing_func(block, x) - if j < len(repmap): - x = repmap[j](x) - x = upscaler(x) - else: - for i, (up_block, upscaler, repmap) in enumerate(block_group): - for j in range(len(repmap) + 1): - for k, block in enumerate(up_block): - if isinstance(block, SDCascadeResBlock): - skip = level_outputs[i] if k == 0 and i > 0 else None - if skip is not None and (x.size(-1) != skip.size(-1) or x.size(-2) != skip.size(-2)): - orig_type = x.dtype - x = torch.nn.functional.interpolate( - x.float(), skip.shape[-2:], mode="bilinear", align_corners=True - ) - x = x.to(orig_type) - x = block(x, skip) - elif isinstance(block, SDCascadeAttnBlock): - x = block(x, clip) - elif isinstance(block, SDCascadeTimestepBlock): - x = block(x, r_embed, delta_r_embed) - else: - x = block(x) - if j < len(repmap): - x = repmap[j](x) - x = upscaler(x) + for i, (up_block, upscaler, repmap) in enumerate(block_group): + for j in range(len(repmap) + 1): + for k, block in enumerate(up_block): + if isinstance(block, SDCascadeResBlock): + skip = level_outputs[i] if k == 0 and i > 0 else None + if skip is not None and (x.size(-1) != skip.size(-1) or x.size(-2) != skip.size(-2)): + orig_type = x.dtype + x = torch.nn.functional.interpolate( + x.float(), skip.shape[-2:], mode="bilinear", align_corners=True + ) + x = x.to(orig_type) + x = self._run_checkpointable_block(checkpoint_block_index, block, x, skip) + elif isinstance(block, SDCascadeAttnBlock): + x = self._run_checkpointable_block(checkpoint_block_index, block, x, clip) + elif isinstance(block, SDCascadeTimestepBlock): + x = self._run_checkpointable_block(checkpoint_block_index, block, x, r_embed, delta_r_embed) + else: + x = self._run_checkpointable_block(checkpoint_block_index, block, x) + checkpoint_block_index += 1 + if j < len(repmap): + x = repmap[j](x) + x = upscaler(x) return x def forward( @@ -665,8 +651,14 @@ def forward( x = x + nn.functional.interpolate( self.pixels_mapper(pixels), size=x.shape[-2:], mode="bilinear", align_corners=True ) - level_outputs = self._down_encode(x, timestep_ratio_embed, clip, delta_timestep_ratio_embed) - x = self._up_decode(level_outputs, timestep_ratio_embed, clip, delta_timestep_ratio_embed) + level_outputs, checkpoint_block_index = self._down_encode(x, timestep_ratio_embed, clip, delta_timestep_ratio_embed) + x = self._up_decode( + level_outputs, + timestep_ratio_embed, + clip, + delta_timestep_ratio_embed, + checkpoint_block_index=checkpoint_block_index, + ) sample = self.clf(x) if not return_dict: diff --git a/simpletuner/helpers/models/wan/model.py b/simpletuner/helpers/models/wan/model.py index 5f29c845c..6b9e92775 100644 --- a/simpletuner/helpers/models/wan/model.py +++ b/simpletuner/helpers/models/wan/model.py @@ -162,6 +162,53 @@ def add_first_frame_conditioning( return conditioned_latent +@torch.no_grad() +def add_first_frame_latent_conditioning( + latent_model_input: torch.Tensor, + clean_latents: torch.Tensor, + vae: AutoencoderKLWan, +): + """ + Adds Wan 2.1 I2V conditioning from already-cached clean latents when the + batch does not carry explicit first-frame pixels. + """ + device = latent_model_input.device + dtype = latent_model_input.dtype + vae_config = getattr(vae, "config", None) + temporal_downsample = getattr(vae_config, "temperal_downsample", None) + if temporal_downsample is None: + temporal_downsample = getattr(vae, "temperal_downsample", None) + vae_scale_factor_temporal = 2 ** sum(temporal_downsample) if temporal_downsample is not None else 4 + + batch, _, num_latent_frames, latent_height, latent_width = latent_model_input.shape + num_frames = (num_latent_frames - 1) * vae_scale_factor_temporal + 1 + + clean_latents = clean_latents.to(device=device, dtype=dtype) + if clean_latents.shape[0] != batch: + clean_latents = clean_latents.expand(batch, -1, -1, -1, -1) + + mask_lat_size = torch.ones( + batch, + 1, + num_frames, + latent_height, + latent_width, + device=device, + dtype=dtype, + ) + mask_lat_size[:, :, 1:] = 0 + first_frame_mask = mask_lat_size[:, :, 0:1] + first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=vae_scale_factor_temporal) + mask_lat_size = torch.concat([first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2) + mask_lat_size = mask_lat_size.view(batch, -1, vae_scale_factor_temporal, latent_height, latent_width) + mask_lat_size = mask_lat_size.transpose(1, 2) + + latent_condition = torch.zeros_like(clean_latents) + latent_condition[:, :, :1] = clean_latents[:, :, :1] + first_frame_condition = torch.concat([mask_lat_size, latent_condition], dim=1) + return torch.cat([latent_model_input, first_frame_condition], dim=1) + + @torch.no_grad() def add_first_frame_conditioning_v22( latent_model_input: torch.Tensor, @@ -758,7 +805,14 @@ def _apply_i2v_conditioning_to_kwargs(self, prepared_batch, transformer_kwargs): return first_frame, last_frame = self._extract_conditioning_frames(prepared_batch) if first_frame is None: - if is_i2v_batch and not getattr(self, "_wan_warned_missing_i2v_conditioning", False) and should_log(): + clean_latents = prepared_batch.get("latents") + if torch.is_tensor(clean_latents): + transformer_kwargs["hidden_states"] = add_first_frame_latent_conditioning( + transformer_kwargs["hidden_states"], + clean_latents, + self.get_vae(), + ) + elif is_i2v_batch and not getattr(self, "_wan_warned_missing_i2v_conditioning", False) and should_log(): logger.warning( "Wan I2V conditioning data was requested but no conditioning frames were provided. " "Ensure your dataset supplies conditioning images when training I2V flavours." diff --git a/simpletuner/helpers/models/wan/transformer.py b/simpletuner/helpers/models/wan/transformer.py index cc6b7ccf8..b55e15554 100644 --- a/simpletuner/helpers/models/wan/transformer.py +++ b/simpletuner/helpers/models/wan/transformer.py @@ -34,7 +34,9 @@ from simpletuner.helpers.models.flowmap import register_flowmap_config from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter @@ -532,6 +534,9 @@ def forward( encoder_hidden_states: torch.Tensor, temb: torch.Tensor, rotary_emb: torch.Tensor, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ) -> torch.Tensor: self._ensure_module_dtype(hidden_states.device, hidden_states.dtype) @@ -557,28 +562,41 @@ def forward( # 1. Self-attention norm_hidden_states = self.norm1(hidden_states) norm_hidden_states = norm_hidden_states * (1 + scale_msa) + shift_msa - attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb) hidden_states = hidden_states + attn_output * gate_msa # 2. Cross-attention norm_hidden_states = self.norm2(hidden_states) - attn_output = self.attn2( - hidden_states=norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - ) + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_output = self.attn2( + hidden_states=norm_hidden_states, + encoder_hidden_states=encoder_hidden_states, + ) hidden_states = hidden_states + attn_output # 3. Feed-forward norm_hidden_states = self.norm3(hidden_states) norm_hidden_states = norm_hidden_states * (1 + c_scale_msa) + c_shift_msa - if self._chunk_enabled: - ff_output = self._run_chunked_feed_forward(norm_hidden_states) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + ff_output = checkpoint_fn( + self._run_feed_forward, + norm_hidden_states, + use_reentrant=False, + ) else: - ff_output = self.ffn(norm_hidden_states) + ff_output = self._run_feed_forward(norm_hidden_states) hidden_states = hidden_states + ff_output * c_gate_msa return hidden_states + def _run_feed_forward(self, norm_hidden_states: torch.Tensor) -> torch.Tensor: + if self._chunk_enabled: + return self._run_chunked_feed_forward(norm_hidden_states) + return self.ffn(norm_hidden_states) + def _run_chunked_feed_forward(self, norm_hidden_states: torch.Tensor) -> torch.Tensor: if self._chunk_auto: return self._auto_chunk_feed_forward(norm_hidden_states) @@ -673,6 +691,8 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOrigi """ _supports_gradient_checkpointing = True + _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _tread_router: Optional[TREADRouter] = None _tread_routes: Optional[List[Dict[str, Any]]] = None _skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"] @@ -803,6 +823,9 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.force_v2_1_time_embedding: bool = False self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, @@ -834,6 +857,15 @@ def set_time_embedding_v2_1(self, force_2_1_time_embedding: bool) -> None: def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): """Set the TREAD router and routing configuration.""" self._tread_router = router @@ -995,107 +1027,180 @@ def _to_pos(idx): grounding_objs = self.position_net(**grounding_kwargs) captured_frame_hidden: Optional[torch.Tensor] = None + use_segmented_checkpointing = ( + torch.is_grad_enabled() + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and skip_layers is None + and not use_routing + and not musubi_offload_active + and grounding_objs is None + and not output_hidden_states + and hidden_states_buffer is None + ) # Transformer blocks with TREAD routing - for i, block in enumerate(self.blocks): - # TREAD: START a route? - if use_routing and route_ptr < len(routes) and i == routes[route_ptr]["start_layer_idx"]: - mask_ratio = routes[route_ptr]["selection_ratio"] - - # Apply routing to video tokens only - # Note: encoder_hidden_states (text) is never routed, only passed to cross-attention - tread_mask_info = router.get_mask( - hidden_states, # (B, S_video, D) where S_video = T*H*W tokens - mask_ratio=mask_ratio, - force_keep=force_keep_mask, - ) - saved_tokens = hidden_states.clone() - hidden_states = router.start_route(hidden_states, tread_mask_info) - routing_now = True - - # Route the rotary embeddings to match the selected video tokens - # This preserves the 3D positional information for kept tokens - current_rope = self._route_rope( - rotary_emb, - tread_mask_info, - keep_len=hidden_states.size(1), - batch=hidden_states.size(0), - ) + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint - # Skip layers if specified - if skip_layers is not None and i in skip_layers: - continue + checkpoint_fn = offloaded_checkpoint + else: + checkpoint_fn = torch.utils.checkpoint.checkpoint - if musubi_offload_active and musubi_manager.is_managed_block(i): - musubi_manager.stream_in(block, hidden_states.device) + def run_wan_block(_block_index, segment_block, segment_hidden_states): + return segment_block( + segment_hidden_states, + encoder_hidden_states, + timestep_proj, + current_rope, + offload_attention=self.gradient_checkpointing_offload_attention, + ) - # Apply transformer block - # Each block does: - # 1. Self-attention on video tokens (with rotary embeddings) - # 2. Cross-attention from video to text tokens - # 3. Feed-forward on video tokens - # Only video tokens are routed; text tokens always remain full sequence - if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": - from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + (hidden_states,) = checkpoint_sequential_state( + self.blocks, + self.gradient_checkpointing_interval, + (hidden_states,), + run_wan_block, + checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + else: + for i, block in enumerate(self.blocks): + # TREAD: START a route? + if use_routing and route_ptr < len(routes) and i == routes[route_ptr]["start_layer_idx"]: + mask_ratio = routes[route_ptr]["selection_ratio"] + + # Apply routing to video tokens only + # Note: encoder_hidden_states (text) is never routed, only passed to cross-attention + tread_mask_info = router.get_mask( + hidden_states, # (B, S_video, D) where S_video = T*H*W tokens + mask_ratio=mask_ratio, + force_keep=force_keep_mask, + ) + saved_tokens = hidden_states.clone() + hidden_states = router.start_route(hidden_states, tread_mask_info) + routing_now = True + + # Route the rotary embeddings to match the selected video tokens + # This preserves the 3D positional information for kept tokens + current_rope = self._route_rope( + rotary_emb, + tread_mask_info, + keep_len=hidden_states.size(1), + batch=hidden_states.size(0), + ) - checkpoint_fn = offloaded_checkpoint + # Skip layers if specified + if skip_layers is not None and i in skip_layers: + continue + + if musubi_offload_active and musubi_manager.is_managed_block(i): + musubi_manager.stream_in(block, hidden_states.device) + + # Apply transformer block + # Each block does: + # 1. Self-attention on video tokens (with rotary embeddings) + # 2. Cross-attention from video to text tokens + # 3. Feed-forward on video tokens + # Only video tokens are routed; text tokens always remain full sequence + if torch.is_grad_enabled() and self.gradient_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + checkpoint_fn = offloaded_checkpoint + else: + checkpoint_fn = torch.utils.checkpoint.checkpoint + + if self.gradient_checkpointing_backend.endswith("-ffn"): + hidden_states = block( + hidden_states, + encoder_hidden_states, + timestep_proj, + current_rope, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_checkpointed_block( + checkpoint_hidden_states, + checkpoint_encoder_hidden_states, + checkpoint_temb, + checkpoint_rope, + checkpoint_block=block, + ): + return checkpoint_block( + checkpoint_hidden_states, + checkpoint_encoder_hidden_states, + checkpoint_temb, + checkpoint_rope, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + hidden_states = checkpoint_fn( + run_checkpointed_block, + hidden_states, # video tokens (possibly routed) + encoder_hidden_states, # text tokens (always full sequence) + timestep_proj, + current_rope, # rotary embeddings (possibly routed) + use_reentrant=False, + ) else: - checkpoint_fn = torch.utils.checkpoint.checkpoint + hidden_states = block( + hidden_states, + encoder_hidden_states, + timestep_proj, + current_rope, + offload_attention=self.gradient_checkpointing_offload_attention, + ) - hidden_states = checkpoint_fn( - block, - hidden_states, # video tokens (possibly routed) - encoder_hidden_states, # text tokens (always full sequence) - timestep_proj, - current_rope, # rotary embeddings (possibly routed) - use_reentrant=False, - ) - else: - hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, current_rope) - - if grounding_objs is not None and hasattr(block, "fuser"): - hidden_states = apply_grounding_fuser( - block.fuser, - hidden_states, - grounding_objs, - tokens_per_frame=post_patch_height * post_patch_width, - num_frames=post_patch_num_frames, - ) + if grounding_objs is not None and hasattr(block, "fuser"): + hidden_states = apply_grounding_fuser( + block.fuser, + hidden_states, + grounding_objs, + tokens_per_frame=post_patch_height * post_patch_width, + num_frames=post_patch_num_frames, + ) - # TREAD: END the current route? - if routing_now and i == routes[route_ptr]["end_layer_idx"]: - hidden_states = router.end_route( - hidden_states, - tread_mask_info, - original_x=saved_tokens, - ) - routing_now = False - route_ptr += 1 - current_rope = rotary_emb - - if output_hidden_states and (hidden_state_layer is None or i == hidden_state_layer): - captured_frame_hidden = hidden_states.reshape( - batch_size, - post_patch_num_frames, - post_patch_height * post_patch_width, - -1, - ) - if hidden_state_layer is not None and i == hidden_state_layer: - output_hidden_states = False - if hidden_states_buffer is not None: - tokens_view = hidden_states.reshape( - batch_size, - post_patch_num_frames, - post_patch_height * post_patch_width, - -1, - ) - else: - tokens_view = hidden_states - _store_hidden_state(hidden_states_buffer, f"layer_{i}", tokens_view) + # TREAD: END the current route? + if routing_now and i == routes[route_ptr]["end_layer_idx"]: + hidden_states = router.end_route( + hidden_states, + tread_mask_info, + original_x=saved_tokens, + ) + routing_now = False + route_ptr += 1 + current_rope = rotary_emb + + if output_hidden_states and (hidden_state_layer is None or i == hidden_state_layer): + captured_frame_hidden = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height * post_patch_width, + -1, + ) + if hidden_state_layer is not None and i == hidden_state_layer: + output_hidden_states = False + if hidden_states_buffer is not None: + tokens_view = hidden_states.reshape( + batch_size, + post_patch_num_frames, + post_patch_height * post_patch_width, + -1, + ) + else: + tokens_view = hidden_states + _store_hidden_state(hidden_states_buffer, f"layer_{i}", tokens_view) - if musubi_offload_active and musubi_manager.is_managed_block(i): - musubi_manager.stream_out(block) + if musubi_offload_active and musubi_manager.is_managed_block(i): + musubi_manager.stream_out(block) # Output processing remains the same if temb.ndim == 2: diff --git a/simpletuner/helpers/models/wan_s2v/model.py b/simpletuner/helpers/models/wan_s2v/model.py index be949034a..d6f3eafc3 100644 --- a/simpletuner/helpers/models/wan_s2v/model.py +++ b/simpletuner/helpers/models/wan_s2v/model.py @@ -6,7 +6,10 @@ import torch import torchaudio from diffusers import AutoencoderKLWan, FlowMatchEulerDiscreteScheduler -from transformers import T5TokenizerFast, UMT5EncoderModel, Wav2Vec2Model, Wav2Vec2Processor +from huggingface_hub import hf_hub_download +from huggingface_hub.errors import EntryNotFoundError, HfHubHTTPError, LocalEntryNotFoundError +from safetensors import SafetensorError, safe_open +from transformers import T5TokenizerFast, UMT5EncoderModel, Wav2Vec2FeatureExtractor, Wav2Vec2Model from simpletuner.helpers.models.common import ModelTypes, PipelineTypes, PredictionTypes, VideoModelFoundation from simpletuner.helpers.models.tae.types import VideoTAESpec @@ -100,7 +103,7 @@ def _load_audio_encoder(self): return logger.info(f"Loading Wav2Vec2 audio encoder from {self.AUDIO_ENCODER_MODEL}") - self._audio_processor = Wav2Vec2Processor.from_pretrained(self.AUDIO_ENCODER_MODEL) + self._audio_processor = Wav2Vec2FeatureExtractor.from_pretrained(self.AUDIO_ENCODER_MODEL) self._audio_encoder = Wav2Vec2Model.from_pretrained( self.AUDIO_ENCODER_MODEL, torch_dtype=torch.float32, # Wav2Vec2 needs fp32 @@ -246,9 +249,12 @@ def _encode_prompts(self, prompts: list, is_negative_prompt: bool = False): if self.config.t5_padding == "zero": prompt_embeds = prompt_embeds * masks.to(device=prompt_embeds.device).unsqueeze(-1).expand(prompt_embeds.shape) - return prompt_embeds, masks + return { + "prompt_embeds": prompt_embeds, + "attention_masks": masks, + } - def prepare_batch_conditions(self, batch: dict) -> dict: + def prepare_batch_conditions(self, batch: dict, state: Optional[dict] = None) -> dict: """ Prepare batch with audio conditioning. @@ -445,11 +451,6 @@ def check_user_config(self): self.config.vae_enable_tiling = True self.config.vae_enable_slicing = True - def pretrained_load_args(self, pretrained_load_args: dict) -> dict: - """Arguments for loading pretrained model.""" - load_args = super().pretrained_load_args(pretrained_load_args) - return load_args - def get_pipeline(self, pipeline_type: PipelineTypes, load_base_model: bool = True): """Get inference pipeline for validation.""" if pipeline_type not in self.PIPELINE_CLASSES: @@ -463,8 +464,14 @@ def get_pipeline(self, pipeline_type: PipelineTypes, load_base_model: bool = Tru else: transformer = None - tokenizer = self.get_tokenizer() - text_encoder = self.get_text_encoder() + if self.tokenizers is None: + self.load_text_tokenizer() + tokenizer = self.tokenizers[0] if self.tokenizers else None + + text_encoder = self.get_text_encoder(0) + if text_encoder is None: + self.load_text_encoder() + text_encoder = self.get_text_encoder(0) vae = self.get_vae() scheduler = FlowMatchEulerDiscreteScheduler() @@ -529,3 +536,106 @@ def setup_model_flavour(self): self.config.pretrained_model_name_or_path = self.HUGGINGFACE_PATHS[flavour] logger.info(f"Configured {self.NAME} with flavour: {flavour}") + + def pretrained_load_args(self, pretrained_load_args: dict) -> dict: + args = super().pretrained_load_args(pretrained_load_args) + args["low_cpu_mem_usage"] = True + return args + + @staticmethod + def _get_parameter(module: torch.nn.Module, name: str) -> torch.nn.Parameter | None: + target = module + for part in name.split(".")[:-1]: + target = getattr(target, part, None) + if target is None: + return None + parameter = getattr(target, name.split(".")[-1], None) + return parameter if isinstance(parameter, torch.nn.Parameter) else None + + @staticmethod + def _set_parameter(module: torch.nn.Module, name: str, tensor: torch.Tensor, requires_grad: bool) -> None: + target = module + parts = name.split(".") + for part in parts[:-1]: + target = getattr(target, part) + setattr(target, parts[-1], torch.nn.Parameter(tensor, requires_grad=requires_grad)) + + @staticmethod + def _checkpoint_file(model_path: str, model_subfolder: str | None, filename: str, load_kwargs: dict) -> str: + relpath = f"{model_subfolder}/{filename}" if model_subfolder else filename + if os.path.isdir(model_path): + return os.path.join(model_path, relpath) + return hf_hub_download( + model_path, + relpath, + revision=load_kwargs.get("revision"), + local_files_only=bool(load_kwargs.get("local_files_only", False)), + ) + + def _load_checkpoint_tensor( + self, + model_path: str, + model_subfolder: str | None, + tensor_name: str, + load_kwargs: dict, + ) -> torch.Tensor | None: + import json + + try: + index_path = self._checkpoint_file( + model_path, + model_subfolder, + "diffusion_pytorch_model.safetensors.index.json", + load_kwargs, + ) + with open(index_path, "r", encoding="utf-8") as handle: + weight_map = json.load(handle).get("weight_map", {}) + except (OSError, json.JSONDecodeError, EntryNotFoundError, HfHubHTTPError, LocalEntryNotFoundError) as exc: + logger.debug("Skipping WanS2V checkpoint alias materialization; could not read safetensors index: %s", exc) + return None + shard_name = weight_map.get(tensor_name) + if shard_name is None: + return None + try: + shard_path = self._checkpoint_file(model_path, model_subfolder, shard_name, load_kwargs) + with safe_open(shard_path, framework="pt", device="cpu") as shard: + return shard.get_tensor(tensor_name) + except (OSError, SafetensorError, KeyError, EntryNotFoundError, HfHubHTTPError, LocalEntryNotFoundError) as exc: + logger.debug("Skipping WanS2V checkpoint alias %s; could not read tensor: %s", tensor_name, exc) + return None + + def materialize_meta_tensors_after_load( + self, + model: torch.nn.Module, + model_path: str, + model_subfolder: str | None, + load_kwargs: dict, + ) -> bool: + aliases = { + "condition_embedder.causal_audio_encoder.weighted_avg.weights": "condition_embedder.causal_audio_encoder.weights", + "condition_embedder.causal_audio_encoder.encoder.conv2.conv.conv.weight": "condition_embedder.causal_audio_encoder.encoder.conv2.conv.weight", + "condition_embedder.causal_audio_encoder.encoder.conv2.conv.conv.bias": "condition_embedder.causal_audio_encoder.encoder.conv2.conv.bias", + "condition_embedder.causal_audio_encoder.encoder.conv3.conv.conv.weight": "condition_embedder.causal_audio_encoder.encoder.conv3.conv.weight", + "condition_embedder.causal_audio_encoder.encoder.conv3.conv.conv.bias": "condition_embedder.causal_audio_encoder.encoder.conv3.conv.bias", + } + materialized = False + for target_name, source_name in aliases.items(): + target_parameter = self._get_parameter(model, target_name) + if target_parameter is None or target_parameter.device.type != "meta": + continue + source_tensor = self._load_checkpoint_tensor(model_path, model_subfolder, source_name, load_kwargs) + if source_tensor is None: + continue + if tuple(source_tensor.shape) != tuple(target_parameter.shape): + logger.warning( + "Skipping WanS2V checkpoint alias %s -> %s because shapes differ: %s != %s", + source_name, + target_name, + tuple(source_tensor.shape), + tuple(target_parameter.shape), + ) + continue + source_tensor = source_tensor.to(dtype=target_parameter.dtype, device="cpu") + self._set_parameter(model, target_name, source_tensor, requires_grad=target_parameter.requires_grad) + materialized = True + return materialized diff --git a/simpletuner/helpers/models/wan_s2v/transformer.py b/simpletuner/helpers/models/wan_s2v/transformer.py index 14a47e9f8..d1951c10b 100644 --- a/simpletuner/helpers/models/wan_s2v/transformer.py +++ b/simpletuner/helpers/models/wan_s2v/transformer.py @@ -40,6 +40,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.tread import TREADRouter logger = logging.get_logger(__name__) @@ -135,10 +136,10 @@ def _apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor) -> torch for i in range(batch_size): s = hidden_states.size(2) x_i = torch.view_as_complex(hidden_states[i, :, :s].to(torch.float64).reshape(hidden_states.size(1), s, -1, 2)) - freqs_i = freqs[i, :s] + freqs_i = freqs[i, :s].transpose(0, 1) x_i = torch.view_as_real(x_i * freqs_i).flatten(2) output.append(x_i) - return torch.stack(output).transpose(1, 2).type_as(hidden_states) + return torch.stack(output).type_as(hidden_states) # ----------------------------------------------------------------------------- @@ -332,7 +333,10 @@ def forward( else: attn_hidden_states = self.injector_pre_norm_feat[audio_attn_id](input_hidden_states) - residual_out = self.injector[audio_attn_id](attn_hidden_states, attn_audio_emb, None, None) + residual_out = self.injector[audio_attn_id]( + hidden_states=attn_hidden_states, + encoder_hidden_states=attn_audio_emb, + ) residual_out = residual_out.unflatten(0, (-1, merged_audio_emb_num_frames)).flatten(1, 2) hidden_states[:, :original_sequence_length] = hidden_states[:, :original_sequence_length] + residual_out @@ -1048,6 +1052,9 @@ def __init__( self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5) self.gradient_checkpointing = False + self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self._musubi_block_swap = MusubiBlockSwapManager.build( depth=num_layers, blocks_to_swap=musubi_blocks_to_swap, @@ -1059,9 +1066,27 @@ def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_t self.condition_embedder.enable_flowmap_time_conditioning(gate_value=gate_value, deltatime_type=deltatime_type) register_flowmap_config(self, gate_value, deltatime_type) - def _set_gradient_checkpointing(self, module, value=False): + def _set_gradient_checkpointing(self, module=None, value=False, enable=None, gradient_checkpointing_func=None): + if enable is not None: + value = enable + if gradient_checkpointing_func is not None: + self._gradient_checkpointing_func = gradient_checkpointing_func + if module is None: + module = self + self.gradient_checkpointing = value if hasattr(module, "gradient_checkpointing"): module.gradient_checkpointing = value + for child in module.children(): + self._set_gradient_checkpointing(child, value=value) + + def set_gradient_checkpointing_backend(self, backend: str): + self.gradient_checkpointing_backend = backend + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): """Set the TREAD router and routing configuration.""" @@ -1308,7 +1333,12 @@ def _to_pos(idx: int) -> int: if musubi_offload_active and musubi_manager.is_managed_block(block_idx): musubi_manager.stream_in(block, hidden_states.device) - if self.training and self.gradient_checkpointing: + if self.training and should_checkpoint_block( + block_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): hidden_states = self._gradient_checkpointing_func( block, hidden_states, encoder_hidden_states, timestep_proj, current_rope ) diff --git a/simpletuner/helpers/models/z_image/transformer.py b/simpletuner/helpers/models/z_image/transformer.py index 2b73c6240..8c396046b 100644 --- a/simpletuner/helpers/models/z_image/transformer.py +++ b/simpletuner/helpers/models/z_image/transformer.py @@ -45,6 +45,8 @@ shard_cp_tensor, unshard_cp_tensor, ) +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state +from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload_context from simpletuner.helpers.training.qk_clip_logging import publish_attention_max_logits from simpletuner.helpers.training.tread import TREADRouter @@ -374,6 +376,9 @@ def forward( attn_mask: torch.Tensor, freqs_cis: torch.Tensor, adaln_input: Optional[torch.Tensor] = None, + checkpoint_ffn: bool = False, + checkpoint_fn: Any | None = None, + offload_attention: bool = False, ): if self.modulation: assert adaln_input is not None @@ -388,45 +393,53 @@ def forward( scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp # Attention block - attn_out = clamp_fp16( - self.attention( - self.attention_norm1(x) * scale_msa, - attention_mask=attn_mask, - freqs_cis=freqs_cis, + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_out = clamp_fp16( + self.attention( + self.attention_norm1(x) * scale_msa, + attention_mask=attn_mask, + freqs_cis=freqs_cis, + ) ) - ) x = x + gate_msa * self.attention_norm2(attn_out) # FFN block - x = x + gate_mlp * self.ffn_norm2( - clamp_fp16( - self.feed_forward( - self.ffn_norm1(x) * scale_mlp, - ) - ) - ) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + ffn_out = checkpoint_fn(self._ffn_forward_modulated, x, scale_mlp, use_reentrant=False) + else: + ffn_out = self._ffn_forward_modulated(x, scale_mlp) + x = x + gate_mlp * ffn_out else: # Attention block - attn_out = clamp_fp16( - self.attention( - self.attention_norm1(x), - attention_mask=attn_mask, - freqs_cis=freqs_cis, + with activation_offload_context(offload_attention, label=f"{self.__class__.__qualname__}:attention"): + attn_out = clamp_fp16( + self.attention( + self.attention_norm1(x), + attention_mask=attn_mask, + freqs_cis=freqs_cis, + ) ) - ) x = x + self.attention_norm2(attn_out) # FFN block - x = x + self.ffn_norm2( - clamp_fp16( - self.feed_forward( - self.ffn_norm1(x), - ) - ) - ) + if checkpoint_ffn: + if checkpoint_fn is None: + raise ValueError("checkpoint_fn is required when checkpoint_ffn=True") + ffn_out = checkpoint_fn(self._ffn_forward_unmodulated, x, use_reentrant=False) + else: + ffn_out = self._ffn_forward_unmodulated(x) + x = x + ffn_out return x + def _ffn_forward_modulated(self, x: torch.Tensor, scale_mlp: torch.Tensor) -> torch.Tensor: + return self.ffn_norm2(clamp_fp16(self.feed_forward(self.ffn_norm1(x) * scale_mlp))) + + def _ffn_forward_unmodulated(self, x: torch.Tensor) -> torch.Tensor: + return self.ffn_norm2(clamp_fp16(self.feed_forward(self.ffn_norm1(x)))) + class FinalLayer(nn.Module): def __init__(self, hidden_size, out_channels): @@ -497,6 +510,8 @@ def __call__(self, ids: torch.Tensor): class ZImageTransformer2DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): _supports_gradient_checkpointing = True + _supports_ffn_gradient_checkpointing = True + _supports_attention_activation_offload = True _no_split_modules = ["ZImageTransformerBlock"] _repeated_blocks = ["ZImageTransformerBlock"] _skip_layerwise_casting_patterns = ["t_embedder", "cap_embedder"] # precision sensitive layers @@ -540,6 +555,9 @@ def __init__( self.t_scale = t_scale self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_offload_attention = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None assert len(all_patch_size) == len(all_f_patch_size) @@ -622,6 +640,15 @@ def __init__( def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_offload_attention(self, enabled: bool): + self.gradient_checkpointing_offload_attention = bool(enabled) + + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: List[Dict[str, Any]]): self._tread_router = router self._tread_routes = routes @@ -855,7 +882,7 @@ def forward( x_attn_mask[i, :seq_len] = 1 if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -863,10 +890,43 @@ def forward( checkpoint_fn = torch.utils.checkpoint.checkpoint for layer in self.noise_refiner: - x = checkpoint_fn(layer, x, x_attn_mask, x_freqs_cis, adaln_input, use_reentrant=False) + if self.gradient_checkpointing_backend.endswith("-ffn"): + x = layer( + x, + x_attn_mask, + x_freqs_cis, + adaln_input, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_noise_refiner( + checkpoint_x, + checkpoint_attn_mask, + checkpoint_freqs, + checkpoint_adaln, + layer_module=layer, + ): + return layer_module( + checkpoint_x, + checkpoint_attn_mask, + checkpoint_freqs, + checkpoint_adaln, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + x = checkpoint_fn(run_noise_refiner, x, x_attn_mask, x_freqs_cis, adaln_input, use_reentrant=False) else: for layer in self.noise_refiner: - x = layer(x, x_attn_mask, x_freqs_cis, adaln_input) + x = layer( + x, + x_attn_mask, + x_freqs_cis, + adaln_input, + offload_attention=self.gradient_checkpointing_offload_attention, + ) # cap embed & refine cap_item_seqlens = [len(_) for _ in cap_feats] @@ -886,7 +946,7 @@ def forward( cap_attn_mask[i, :seq_len] = 1 if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint @@ -894,10 +954,41 @@ def forward( checkpoint_fn = torch.utils.checkpoint.checkpoint for layer in self.context_refiner: - cap_feats = checkpoint_fn(layer, cap_feats, cap_attn_mask, cap_freqs_cis, use_reentrant=False) + if self.gradient_checkpointing_backend.endswith("-ffn"): + cap_feats = layer( + cap_feats, + cap_attn_mask, + cap_freqs_cis, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + else: + + def run_context_refiner( + checkpoint_cap_feats, + checkpoint_attn_mask, + checkpoint_freqs, + layer_module=layer, + ): + return layer_module( + checkpoint_cap_feats, + checkpoint_attn_mask, + checkpoint_freqs, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + cap_feats = checkpoint_fn( + run_context_refiner, cap_feats, cap_attn_mask, cap_freqs_cis, use_reentrant=False + ) else: for layer in self.context_refiner: - cap_feats = layer(cap_feats, cap_attn_mask, cap_freqs_cis) + cap_feats = layer( + cap_feats, + cap_attn_mask, + cap_freqs_cis, + offload_attention=self.gradient_checkpointing_offload_attention, + ) # unified unified = [] @@ -961,15 +1052,41 @@ def _to_pos(idx): def apply_layer(layer_module, h, attn_mask, freqs, layer_adaln_input): if torch.is_grad_enabled() and self.gradient_checkpointing: - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint else: checkpoint_fn = torch.utils.checkpoint.checkpoint - return checkpoint_fn(layer_module, h, attn_mask, freqs, layer_adaln_input, use_reentrant=False) - return layer_module(h, attn_mask, freqs, layer_adaln_input) + if self.gradient_checkpointing_backend.endswith("-ffn"): + return layer_module( + h, + attn_mask, + freqs, + layer_adaln_input, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + def run_checkpointed_layer(checkpoint_h, checkpoint_attn_mask, checkpoint_freqs, checkpoint_adaln): + return layer_module( + checkpoint_h, + checkpoint_attn_mask, + checkpoint_freqs, + checkpoint_adaln, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + return checkpoint_fn(run_checkpointed_layer, h, attn_mask, freqs, layer_adaln_input, use_reentrant=False) + return layer_module( + h, + attn_mask, + freqs, + layer_adaln_input, + offload_attention=self.gradient_checkpointing_offload_attention, + ) skip_set = set(skip_layers) if skip_layers is not None else set() capture_idx = 0 @@ -982,7 +1099,49 @@ def apply_layer(layer_module, h, attn_mask, freqs, layer_adaln_input): if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, unified.device, grad_enabled) + use_segmented_checkpointing = ( + grad_enabled + and self.gradient_checkpointing + and self.gradient_checkpointing_interval is not None + and self.gradient_checkpointing_interval > 1 + and not self.gradient_checkpointing_backend.endswith("-ffn") + and not use_routing + and not skip_set + and hidden_states_buffer is None + and not musubi_offload_active + ) + if use_segmented_checkpointing: + if self.gradient_checkpointing_backend.startswith("unsloth"): + from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint + + segmented_checkpoint_fn = offloaded_checkpoint + else: + segmented_checkpoint_fn = torch.utils.checkpoint.checkpoint + + layer_adaln_input = unified_adaln_input if unified_adaln_input is not None else adaln_input + + def run_segmented_layer(_idx, layer_module, h): + return layer_module( + h, + unified_attn_mask, + unified_freqs_cis, + layer_adaln_input, + offload_attention=self.gradient_checkpointing_offload_attention, + ) + + (unified,) = checkpoint_sequential_state( + self.layers, + self.gradient_checkpointing_interval, + (unified,), + run_segmented_layer, + segmented_checkpoint_fn, + {"use_reentrant": False}, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + for idx, layer in enumerate(self.layers): + if use_segmented_checkpointing: + break if musubi_offload_active and musubi_manager.is_managed_block(idx): musubi_manager.stream_in(layer, unified.device) if use_routing and route_ptr < len(routes) and idx == routes[route_ptr]["start_layer_idx"]: diff --git a/simpletuner/helpers/models/zlab_i1/transformer.py b/simpletuner/helpers/models/zlab_i1/transformer.py index e850a03ed..68952a390 100644 --- a/simpletuner/helpers/models/zlab_i1/transformer.py +++ b/simpletuner/helpers/models/zlab_i1/transformer.py @@ -17,6 +17,7 @@ from huggingface_hub import hf_hub_download from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state, should_checkpoint_block from simpletuner.helpers.training.tread import TREADRouter logger = logging.getLogger(__name__) @@ -501,6 +502,8 @@ def __init__( self.register_to_config(head_dim=resolved_head_dim) self.head_dim = resolved_head_dim self.gradient_checkpointing = False + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.input_size = input_size self.image_resolution = image_resolution self.patch_size = patch_size @@ -584,6 +587,12 @@ def __init__( logger=logger, ) + def set_gradient_checkpointing_interval(self, interval: int): + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_router(self, router: TREADRouter, routes: list[dict[str, Any]]): self._tread_router = router self._tread_routes = routes @@ -764,6 +773,48 @@ def full_image_tokens_for_storage(): return router.end_route(image_tokens, tread_mask_info, original_x=saved_image_tokens) return image_tokens + def run_segmented_block( + block_index: int, + block: i1DiTBlock, + segment_image_tokens: torch.Tensor, + segment_text_tokens: torch.Tensor, + *skip_slots: torch.Tensor, + ) -> tuple[torch.Tensor, ...]: + skip_images = list(skip_slots[: len(self.in_blocks)]) + skip_texts = list(skip_slots[len(self.in_blocks) :]) + if block_index < len(self.in_blocks): + segment_image_tokens, segment_text_tokens = block( + segment_image_tokens, + segment_text_tokens, + image_freqs, + text_freqs, + text_mask, + None, + ) + skip_images[block_index] = segment_image_tokens + skip_texts[block_index] = segment_text_tokens + elif block_index == len(self.in_blocks): + segment_image_tokens, segment_text_tokens = block( + segment_image_tokens, + segment_text_tokens, + image_freqs, + text_freqs, + text_mask, + None, + ) + else: + out_index = block_index - len(self.in_blocks) - 1 + skip_index = len(self.in_blocks) - 1 - out_index + segment_image_tokens, segment_text_tokens = block( + segment_image_tokens, + segment_text_tokens, + image_freqs, + text_freqs, + text_mask, + (skip_images[skip_index], skip_texts[skip_index]), + ) + return (segment_image_tokens, segment_text_tokens, *skip_images, *skip_texts) + def run_block(block: i1DiTBlock, skip: Optional[tuple[torch.Tensor, torch.Tensor]] = None): block_skip = skip if ( @@ -773,7 +824,12 @@ def run_block(block: i1DiTBlock, skip: Optional[tuple[torch.Tensor, torch.Tensor and block_skip[0].shape[1] != image_tokens.shape[1] ): block_skip = (router.start_route(block_skip[0], tread_mask_info), block_skip[1]) - if torch.is_grad_enabled() and self.gradient_checkpointing: + if torch.is_grad_enabled() and should_checkpoint_block( + global_idx, + self.gradient_checkpointing, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): inputs = ( block, image_tokens, @@ -793,45 +849,75 @@ def run_block(block: i1DiTBlock, skip: Optional[tuple[torch.Tensor, torch.Tensor if musubi_manager is not None: musubi_offload_active = musubi_manager.activate(combined_blocks, x.device, torch.is_grad_enabled()) - for block in self.in_blocks: - if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_in(block, x.device) - maybe_start_route() - if global_idx not in skip_set: - image_tokens, text_tokens = run_block(block) - image_tokens_for_storage = full_image_tokens_for_storage() - skips.append((image_tokens_for_storage, text_tokens)) - maybe_end_route() - _store_hidden_state(hidden_states_buffer, f"layer_{global_idx}", image_tokens_for_storage) - if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_out(block) - global_idx += 1 + segment_size = self.gradient_checkpointing_interval + use_segmented_checkpointing = ( + self.training + and torch.is_grad_enabled() + and self.gradient_checkpointing + and segment_size is not None + and segment_size > 1 + and not use_routing + and not musubi_offload_active + and skip_layers is None + and hidden_states_buffer is None + ) + if use_segmented_checkpointing: + empty_image_skip = image_tokens.new_empty(0) + empty_text_skip = text_tokens.new_empty(0) + segmented_state = checkpoint_sequential_state( + combined_blocks, + segment_size, + ( + image_tokens, + text_tokens, + *([empty_image_skip] * len(self.in_blocks)), + *([empty_text_skip] * len(self.in_blocks)), + ), + run_segmented_block, + self._gradient_checkpointing_func, + segment_stride=self.gradient_checkpointing_segment_stride, + ) + image_tokens, text_tokens = segmented_state[:2] + else: + for block in self.in_blocks: + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): + musubi_manager.stream_in(block, x.device) + maybe_start_route() + if global_idx not in skip_set: + image_tokens, text_tokens = run_block(block) + image_tokens_for_storage = full_image_tokens_for_storage() + skips.append((image_tokens_for_storage, text_tokens)) + maybe_end_route() + _store_hidden_state(hidden_states_buffer, f"layer_{global_idx}", image_tokens_for_storage) + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): + musubi_manager.stream_out(block) + global_idx += 1 - if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_in(self.mid_block, x.device) - maybe_start_route() - if global_idx not in skip_set: - image_tokens, text_tokens = run_block(self.mid_block) - image_tokens_for_storage = full_image_tokens_for_storage() - maybe_end_route() - _store_hidden_state(hidden_states_buffer, f"layer_{global_idx}", image_tokens_for_storage) - if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_out(self.mid_block) - global_idx += 1 - - for block in self.out_blocks: if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_in(block, x.device) + musubi_manager.stream_in(self.mid_block, x.device) maybe_start_route() - skip = skips.pop() if global_idx not in skip_set: - image_tokens, text_tokens = run_block(block, skip) + image_tokens, text_tokens = run_block(self.mid_block) image_tokens_for_storage = full_image_tokens_for_storage() maybe_end_route() _store_hidden_state(hidden_states_buffer, f"layer_{global_idx}", image_tokens_for_storage) if musubi_offload_active and musubi_manager.is_managed_block(global_idx): - musubi_manager.stream_out(block) + musubi_manager.stream_out(self.mid_block) global_idx += 1 + + for block in self.out_blocks: + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): + musubi_manager.stream_in(block, x.device) + maybe_start_route() + skip = skips.pop() + if global_idx not in skip_set: + image_tokens, text_tokens = run_block(block, skip) + image_tokens_for_storage = full_image_tokens_for_storage() + maybe_end_route() + _store_hidden_state(hidden_states_buffer, f"layer_{global_idx}", image_tokens_for_storage) + if musubi_offload_active and musubi_manager.is_managed_block(global_idx): + musubi_manager.stream_out(block) + global_idx += 1 tokens = self.final_layer(image_tokens) bsz = x.shape[0] h = token_height diff --git a/simpletuner/helpers/multiaspect/sampler.py b/simpletuner/helpers/multiaspect/sampler.py index 2a34480a1..af90a05f3 100644 --- a/simpletuner/helpers/multiaspect/sampler.py +++ b/simpletuner/helpers/multiaspect/sampler.py @@ -9,7 +9,7 @@ from simpletuner.helpers.data_backend.base import BaseDataBackend from simpletuner.helpers.image_manipulation.training_sample import TrainingSample -from simpletuner.helpers.metadata.backends.base import MetadataBackend +from simpletuner.helpers.metadata.backends.base import MetadataBackend, get_cp_aware_dp_info from simpletuner.helpers.multiaspect.image import MultiaspectImage from simpletuner.helpers.multiaspect.state import BucketStateManager from simpletuner.helpers.prompts import PromptHandler @@ -121,6 +121,7 @@ def save_state(self, state_path: str): This method should be called when the accelerator save hook is called, so that the state is correctly restored with a given checkpoint. """ + effective_dp_size, dp_rank, _cp_size = get_cp_aware_dp_info(self.accelerator) state = { "aspect_ratio_bucket_indices": self.metadata_backend.aspect_ratio_bucket_indices, "buckets": self.buckets, @@ -129,15 +130,62 @@ def save_state(self, state_path: str): "current_bucket": self.current_bucket, "seen_images": self.metadata_backend.seen_images, "current_epoch": self.current_epoch, + "dp_size": effective_dp_size, + "dp_rank": dp_rank, } self.state_manager.save_state(state, state_path) + def _saved_schedule_is_restorable(self, previous_state: dict, state_path: str) -> bool: + """A saved schedule is one rank's shard of the dataset. + + Restoring it is only correct when the data-parallel layout is unchanged and every rank + has a shard to restore. Restoring on some ranks while the others re-split leaves the + ranks disagreeing about who owns which samples. + """ + effective_dp_size, dp_rank, _cp_size = get_cp_aware_dp_info(self.accelerator) + if previous_state.get("dp_size") != effective_dp_size or previous_state.get("dp_rank") != dp_rank: + self.logger.warning( + f"Checkpoint schedule was written for dp_size={previous_state.get('dp_size')} " + f"dp_rank={previous_state.get('dp_rank')}, but this run is dp_size={effective_dp_size} " + f"dp_rank={dp_rank}. Keeping the freshly split schedule." + ) + return False + + checkpoint_dir = os.path.dirname(state_path) + own_filename = os.path.basename(state_path) + stem, extension = os.path.splitext(own_filename) + stem = stem.split("-rank", 1)[0] + for rank in range(self.accelerator.num_processes): + # Rank 0 writes the unsuffixed name, so it has to be built separately rather than + # scanned as a -rank{N} sibling. This rank's own state is evidently present. + filename = f"{stem}{extension}" if rank == 0 else f"{stem}-rank{rank}{extension}" + if filename == own_filename: + continue + sibling = self.state_manager.mangle_state_path(os.path.join(checkpoint_dir, filename)) + if not os.path.exists(sibling): + self.logger.warning( + f"Checkpoint holds no sampler state for rank {rank}, so the rank-local schedules " + "cannot be restored consistently. Keeping the freshly split schedule." + ) + return False + return True + def load_states(self, state_path: str): try: - self.buckets = self.load_buckets() previous_state = self.state_manager.load_state(state_path) except Exception as e: raise e + + # Checkpoints contain the rank-local schedule. Restore it before seen + # state so legacy boolean flags can be expanded to all occurrences. + saved_schedule = previous_state.get("aspect_ratio_bucket_indices") + if isinstance(saved_schedule, dict) and self._saved_schedule_is_restorable(previous_state, state_path): + self.metadata_backend.aspect_ratio_bucket_indices = saved_schedule + self._val_master_list = sorted(sum(saved_schedule.values(), [])) + self.buckets = previous_state.get("buckets", self.load_buckets()) + if "current_bucket" in previous_state: + self.current_bucket = previous_state["current_bucket"] + self.exhausted_buckets = [] if "exhausted_buckets" in previous_state: self.logger.info(f"Previous checkpoint had {len(previous_state['exhausted_buckets'])} exhausted buckets.") @@ -149,7 +197,20 @@ def load_states(self, state_path: str): # Merge seen_images into self.state_manager.seen_images Manager.dict: if "seen_images" in previous_state: self.logger.info(f"Previous checkpoint had {len(previous_state['seen_images'])} seen {self.sample_type_strs}.") - self.metadata_backend.seen_images.update(previous_state["seen_images"]) + occurrence_counts = {} + for images in self.metadata_backend.aspect_ratio_bucket_indices.values(): + for image_path in images: + occurrence_counts[image_path] = occurrence_counts.get(image_path, 0) + 1 + normalized_seen = { + image_path: ( + (occurrence_counts.get(image_path, True) if value else 0) + if isinstance(value, bool) + else value + ) + for image_path, value in previous_state["seen_images"].items() + } + self.metadata_backend.seen_images.clear() + self.metadata_backend.seen_images.update(normalized_seen) def load_buckets(self): return list(self.metadata_backend.aspect_ratio_bucket_indices.keys()) # These keys are a float value, eg. 1.78. @@ -389,6 +450,17 @@ def _get_bucket_images(self, bucket): # Bucket not found with either type return [] + def _filter_unseen_occurrences(self, images): + """Filter consumed positions without collapsing duplicate filepaths.""" + occurrence_indices = {} + unseen = [] + for image in images: + occurrence_index = occurrence_indices.get(image, 0) + occurrence_indices[image] = occurrence_index + 1 + if not self.metadata_backend.is_seen(image, occurrence_index): + unseen.append(image) + return unseen + def _get_unseen_images(self, bucket=None): """ Get unseen {self.sample_type_strs} from the specified bucket. @@ -410,8 +482,7 @@ def _get_unseen_images(self, bucket=None): return [ (os.path.join(self.metadata_backend.instance_data_dir, image) if not image.startswith("http") else image) - for image in bucket_images - if not self.metadata_backend.is_seen(image) + for image in self._filter_unseen_occurrences(bucket_images) ] elif bucket is None: unseen_images = [] @@ -423,8 +494,7 @@ def _get_unseen_images(self, bucket=None): if not image.startswith("http") else image ) - for image in images - if not self.metadata_backend.is_seen(image) + for image in self._filter_unseen_occurrences(images) ] ) return unseen_images diff --git a/simpletuner/helpers/musubi_block_swap.py b/simpletuner/helpers/musubi_block_swap.py index fba81e601..a467e93df 100644 --- a/simpletuner/helpers/musubi_block_swap.py +++ b/simpletuner/helpers/musubi_block_swap.py @@ -31,10 +31,28 @@ def _is_quanto_tensor(tensor) -> bool: return module_name.startswith("optimum.quanto.") and hasattr(tensor, "_data") +def _is_sdnq_tensor(tensor) -> bool: + return ( + type(tensor).__module__.startswith("sdnq.") + or type(tensor).__name__ == "SDNQTensor" + or (hasattr(tensor, "sdnq_dequantizer") and hasattr(tensor, "weight") and hasattr(tensor, "scale")) + ) + + +def _is_sdnq_module(module: nn.Module) -> bool: + return type(module).__module__.startswith("sdnq.") or type(module).__name__.startswith("SDNQ") + + def _tensor_on_device(tensor, device: torch.device) -> bool: if not _same_device(tensor.device, device): return False if not _is_quanto_tensor(tensor): + if not _is_sdnq_tensor(tensor): + return True + for attr in ("weight", "scale", "zero_point", "svd_up", "svd_down"): + value = getattr(tensor, attr, None) + if value is not None and hasattr(value, "device") and not _same_device(value.device, device): + return False return True for attr in ("_data", "_scale", "_shift", "_scale_shift"): value = getattr(tensor, attr, None) @@ -55,6 +73,32 @@ def _module_has_quanto_tensor(module: nn.Module) -> bool: ) +def _module_has_sdnq_payload(module: nn.Module) -> bool: + return ( + any(_is_sdnq_module(child) for child in module.modules()) + or any(_is_sdnq_tensor(tensor) for tensor in module.parameters()) + or any(_is_sdnq_tensor(tensor) for tensor in module.buffers()) + ) + + +def _module_has_trainable_local_state(module: nn.Module) -> bool: + return any(param is not None and param.requires_grad for param in module._parameters.values()) + + +def _module_has_local_quantized_payload(module: nn.Module) -> bool: + return ( + any( + param is not None and (_is_quanto_tensor(param) or _is_sdnq_tensor(param)) + for param in module._parameters.values() + ) + or any( + buffer is not None and (_is_quanto_tensor(buffer) or _is_sdnq_tensor(buffer)) + for buffer in module._buffers.values() + ) + or _is_sdnq_module(module) + ) + + def _move_quanto_tensor_to_device(tensor, device: torch.device): if not _same_device(tensor.device, device): tensor.data = tensor.data.to(device, non_blocking=True) @@ -69,12 +113,28 @@ def _move_quanto_tensor_to_device(tensor, device: torch.device): setattr(tensor, attr, value.to(device, non_blocking=True)) -def _move_module_without_swapping_quanto_params(module: nn.Module, device: torch.device): +def _move_sdnq_tensor_to_device(tensor, device: torch.device): + moved = tensor.to(device, non_blocking=True) + for attr in ("weight", "scale", "zero_point", "svd_up", "svd_down"): + value = getattr(moved, attr, None) + if value is None: + value = getattr(tensor, attr, None) + if value is not None and hasattr(value, "device") and not _same_device(value.device, device): + value = value.to(device, non_blocking=True) + if value is not None and hasattr(value, "device"): + setattr(tensor, attr, value) + if not _same_device(tensor.device, device): + tensor.data = moved.data + + +def _move_module_without_swapping_quantized_params(module: nn.Module, device: torch.device): for child in module.children(): - _move_module_without_swapping_quanto_params(child, device) + _move_module_without_swapping_quantized_params(child, device) - keep_local_trainable_state = device.type == "cpu" and any( - param is not None and param.requires_grad for param in module._parameters.values() + keep_local_trainable_state = ( + device.type == "cpu" + and not _module_has_local_quantized_payload(module) + and any(param is not None and param.requires_grad for param in module._parameters.values()) ) for key, param in module._parameters.items(): @@ -84,6 +144,8 @@ def _move_module_without_swapping_quanto_params(module: nn.Module, device: torch continue if _is_quanto_tensor(param): _move_quanto_tensor_to_device(param, device) + elif _is_sdnq_tensor(param): + _move_sdnq_tensor_to_device(param, device) elif not _same_device(param.device, device): param.data = param.data.to(device, non_blocking=True) if param.grad is not None and not _same_device(param.grad.device, device): @@ -96,6 +158,8 @@ def _move_module_without_swapping_quanto_params(module: nn.Module, device: torch continue if _is_quanto_tensor(buffer): _move_quanto_tensor_to_device(buffer, device) + elif _is_sdnq_tensor(buffer): + _move_sdnq_tensor_to_device(buffer, device) elif not _same_device(buffer.device, device): module._buffers[key] = buffer.to(device, non_blocking=True) @@ -221,8 +285,10 @@ def _move_module(self, module: nn.Module, device: torch.device): if _module_on_device(module, device): return with torch.no_grad(): - if _module_has_quanto_tensor(module): - _move_module_without_swapping_quanto_params(module, device) + if _module_has_quanto_tensor(module) or _module_has_sdnq_payload(module): + _move_module_without_swapping_quantized_params(module, device) + elif device.type == "cpu" and any(_module_has_trainable_local_state(child) for child in module.modules()): + _move_module_without_swapping_quantized_params(module, device) else: module.to(device) diff --git a/simpletuner/helpers/prompts.py b/simpletuner/helpers/prompts.py index ac04b7bec..9f156a290 100644 --- a/simpletuner/helpers/prompts.py +++ b/simpletuner/helpers/prompts.py @@ -588,9 +588,21 @@ def get_all_captions( captions = [] caption_image_paths = [] images_missing_captions = [] - all_image_files = StateTracker.get_image_files(data_backend_id=data_backend.id) or data_backend.list_files( - instance_data_dir=instance_data_dir, file_extensions=image_file_extensions - ) + try: + data_backend_state = StateTracker.get_data_backend(data_backend.id) or {} + except KeyError: + data_backend_state = {} + metadata_backend = data_backend_state.get("metadata_backend") + bucket_indices = getattr(metadata_backend, "aspect_ratio_bucket_indices", None) + max_num_samples = getattr(metadata_backend, "max_num_samples", None) or backend_config.get("max_num_samples") + if max_num_samples and isinstance(bucket_indices, dict) and bucket_indices: + all_image_files = [] + for bucket in sorted(bucket_indices.keys(), key=str): + all_image_files.extend(bucket_indices[bucket]) + else: + all_image_files = StateTracker.get_image_files(data_backend_id=data_backend.id) or data_backend.list_files( + instance_data_dir=instance_data_dir, file_extensions=image_file_extensions + ) if isinstance(all_image_files, list) and len(all_image_files) > 0 and isinstance(all_image_files[0], tuple): all_image_files = all_image_files[0][2] from tqdm import tqdm diff --git a/simpletuner/helpers/publishing/huggingface.py b/simpletuner/helpers/publishing/huggingface.py index 56120b1ec..52320b208 100644 --- a/simpletuner/helpers/publishing/huggingface.py +++ b/simpletuner/helpers/publishing/huggingface.py @@ -296,9 +296,14 @@ def find_latest_checkpoint(self): return highest_checkpoint def upload_latest_checkpoint( - self, validation_images: dict, webhook_handler=None, global_step: int = None, epoch: int = None + self, + validation_images: dict, + webhook_handler=None, + global_step: int = None, + epoch: int = None, + checkpoint_path: str = None, ): - checkpoint_path = self.find_latest_checkpoint() + checkpoint_path = Path(checkpoint_path) if checkpoint_path else self.find_latest_checkpoint() if checkpoint_path: logging.info(f"Checkpoint path: {checkpoint_path}") try: diff --git a/simpletuner/helpers/ramtorch/modules/linear.py b/simpletuner/helpers/ramtorch/modules/linear.py index 7d843c575..0ac8e517c 100644 --- a/simpletuner/helpers/ramtorch/modules/linear.py +++ b/simpletuner/helpers/ramtorch/modules/linear.py @@ -9,6 +9,7 @@ - Scenarios where GPU memory is limited but CPU memory is abundant """ +import os from collections import OrderedDict import torch @@ -22,6 +23,20 @@ _DEVICE_STATE = {} +def _env_int(name: str, default: int) -> int: + value = os.environ.get(name) + if value in (None, ""): + return default + try: + return int(value) + except ValueError: + return default + + +def _forward_prefetch_stream_count() -> int: + return max(_env_int("SIMPLETUNER_RAMTORCH_FORWARD_PREFETCH_STREAMS", 4), 1) + + def _to_cpu_pinned(tensor: torch.Tensor, *, dtype: torch.dtype | None = None) -> torch.Tensor: if dtype is not None and tensor.dtype != dtype: tensor = tensor.to(dtype=dtype) @@ -52,6 +67,10 @@ def _get_device_state(device=None): _DEVICE_STATE[device] = { # streams & events "transfer_stream": torch.cuda.Stream(device=device), + "forward_transfer_streams": [ + torch.cuda.Stream(device=device) for _ in range(_forward_prefetch_stream_count()) + ], + "forward_transfer_stream_clk": 0, "transfer_grad_stream": torch.cuda.Stream(device=device), "transfer_backward_finished_event": torch.cuda.Event(), "transfer_weight_backward_start_event": torch.cuda.Event(), @@ -78,6 +97,16 @@ def _get_device_state(device=None): return _DEVICE_STATE[device] +def _next_forward_transfer_stream(state, device): + streams = state.get("forward_transfer_streams") + if not streams: + streams = [state.get("transfer_stream") or torch.cuda.Stream(device=device)] + state["forward_transfer_streams"] = streams + selected = int(state.get("forward_transfer_stream_clk", 0)) % len(streams) + state["forward_transfer_stream_clk"] = selected + 1 + return streams[selected] + + def _prefetch_key(weight_cpu, bias_cpu): return (id(weight_cpu), id(bias_cpu) if bias_cpu is not None else None) @@ -94,6 +123,11 @@ def _resident_bytes(entry): return int(entry.get("bytes", 0)) +def _clear_linear_forward_residency(weight_cpu) -> None: + if hasattr(weight_cpu, "_ramtorch_last_forward_residency"): + delattr(weight_cpu, "_ramtorch_last_forward_residency") + + def _dense_weight_for_grad(weight, *, device=None, dtype=None): if hasattr(weight, "dequantize"): weight = weight.dequantize() @@ -104,6 +138,24 @@ def _dense_weight_for_grad(weight, *, device=None, dtype=None): return weight +def _dense_weight_for_linear(weight, *, dtype=None): + if hasattr(weight, "dequantize"): + weight = weight.dequantize() + if dtype is not None and getattr(weight, "dtype", None) != dtype: + weight = weight.to(dtype) + return weight + + +def _record_stream_if_supported(tensor: torch.Tensor | None, stream) -> None: + if tensor is None: + return + try: + tensor.record_stream(stream) + except (AssertionError, NotImplementedError) as exc: + if "record_stream" not in str(exc): + raise + + def _evict_backward_resident_entry(state, key, device): entry = state["forward_backward_residency"].pop(key, None) if entry is None: @@ -160,32 +212,21 @@ def prefetch_linear_forward(weight_cpu, bias_cpu, device="cuda"): if not allowed: return False - selected_buffer = state["forward_clk"] - state["forward_clk"] ^= 1 - state["forward_buffer_generations"][selected_buffer] += 1 - buffer_generation = state["forward_buffer_generations"][selected_buffer] - - transfer_stream = state["transfer_stream"] - release_event = state["forward_buffer_release_events"][selected_buffer] + transfer_stream = _next_forward_transfer_stream(state, device_obj) with torch.cuda.stream(transfer_stream): - if release_event is not None: - transfer_stream.wait_event(release_event) - with record_function("forward_weight_bias_prefetch"): - state["w_buffers"][selected_buffer] = weight_cpu.to(device_obj, non_blocking=True) - state["b_buffers"][selected_buffer] = ( - bias_cpu.to(device_obj, non_blocking=True) if bias_cpu is not None else None - ) + weight_gpu = weight_cpu.to(device_obj, non_blocking=True) + bias_gpu = bias_cpu.to(device_obj, non_blocking=True) if bias_cpu is not None else None event = torch.cuda.Event() event.record() state["forward_prefetches"][key] = { - "buffer": selected_buffer, "event": event, - "generation": buffer_generation, "versions": versions, + "weight": weight_gpu, + "bias": bias_gpu, } ramtorch_profile.record_prefetch_enqueued( "linear", @@ -207,7 +248,11 @@ def _consume_linear_forward_prefetch(weight_cpu, bias_cpu, device): if entry["versions"] != _prefetch_versions(weight_cpu, bias_cpu): ramtorch_profile.record_prefetch_stale("linear", key, device) return None - if state["forward_buffer_generations"][entry["buffer"]] != entry.get("generation"): + buffer_index = entry.get("buffer") + if buffer_index is not None and state["forward_buffer_generations"][buffer_index] != entry.get("generation"): + ramtorch_profile.record_prefetch_stale("linear", key, device) + return None + if buffer_index is None and entry.get("weight") is None: ramtorch_profile.record_prefetch_stale("linear", key, device) return None ramtorch_profile.record_prefetch_consumed("linear", key, device) @@ -270,11 +315,12 @@ def preserve_linear_forward_for_backward( state = _get_device_state(device_obj) buffer_index = last_forward.get("buffer") - if ( - buffer_index is None - or state["forward_buffer_generations"][buffer_index] != last_forward.get("generation") - or last_forward.get("weight") is None - ): + if last_forward.get("weight") is None: + _clear_linear_forward_residency(weight_cpu) + ramtorch_profile.record_backward_preserve_attempt(device_obj, (weight_cpu, bias_cpu), retained=False) + return False + if buffer_index is not None and state["forward_buffer_generations"][buffer_index] != last_forward.get("generation"): + _clear_linear_forward_residency(weight_cpu) ramtorch_profile.record_backward_preserve_attempt(device_obj, (weight_cpu, bias_cpu), retained=False) return False @@ -306,6 +352,7 @@ def preserve_linear_forward_for_backward( max_bytes=max(int(max_bytes), 0), ) retained = key in residency + _clear_linear_forward_residency(weight_cpu) ramtorch_profile.record_backward_preserve_attempt( device_obj, (weight_cpu, bias_cpu), @@ -469,35 +516,45 @@ def forward(ctx, x, weight_cpu, bias_cpu, device="cuda"): transfer = _consume_linear_forward_prefetch(weight_cpu, bias_cpu, device_obj) if transfer is None: transfer = _transfer_linear_forward(weight_cpu, bias_cpu, device_obj) - selected_buffer = transfer["buffer"] + selected_buffer = transfer.get("buffer") + weight_forward = transfer.get("weight") + bias_forward = transfer.get("bias") + if weight_forward is None: + weight_forward = w_buffers[selected_buffer] + bias_forward = b_buffers[selected_buffer] with torch.cuda.device(device_obj): compute_stream = torch.cuda.current_stream(device_obj) with record_function("forward_linear_compute"): # for profiling and easy debugging # make compute stream wait for this transfer compute_stream.wait_event(transfer["event"]) + _record_stream_if_supported(weight_forward, compute_stream) + if bias_forward is not None: + _record_stream_if_supported(bias_forward, compute_stream) # Manual casting when autocast is enabled if autocast_enabled: x_compute = x.to(autocast_dtype) - w_compute = w_buffers[selected_buffer].to(autocast_dtype) - b_compute = ( - b_buffers[selected_buffer].to(autocast_dtype) if b_buffers[selected_buffer] is not None else None - ) + w_compute = _dense_weight_for_linear(weight_forward, dtype=autocast_dtype) + b_compute = bias_forward.to(autocast_dtype) if bias_forward is not None else None out = F.linear(x_compute, w_compute, b_compute) else: - out = F.linear(x, w_buffers[selected_buffer], b_buffers[selected_buffer]) + w_compute = _dense_weight_for_linear(weight_forward, dtype=x.dtype) + out = F.linear(x, w_compute, bias_forward) release_event = torch.cuda.Event() release_event.record(compute_stream) - state["forward_buffer_release_events"][selected_buffer] = release_event + if selected_buffer is not None: + state["forward_buffer_release_events"][selected_buffer] = release_event weight_cpu._ramtorch_last_forward_residency = { "key": _prefetch_key(weight_cpu, bias_cpu), "versions": _prefetch_versions(weight_cpu, bias_cpu), "buffer": selected_buffer, - "generation": state["forward_buffer_generations"][selected_buffer], + "generation": ( + state["forward_buffer_generations"][selected_buffer] if selected_buffer is not None else None + ), "event": release_event, - "weight": w_buffers[selected_buffer], - "bias": b_buffers[selected_buffer], + "weight": weight_forward, + "bias": bias_forward, } # save for backward @@ -894,5 +951,9 @@ def preserve_forward_for_backward(self, *, max_entries: int = 2, max_bytes: int max_bytes=max_bytes, ) + def discard_forward_for_backward(self): + """Release the most recent forward GPU weight reference when it will not be reused.""" + _clear_linear_forward_residency(self.weight) + Linear = CPUBouncingLinear diff --git a/simpletuner/helpers/ramtorch/profiling.py b/simpletuner/helpers/ramtorch/profiling.py index a03ba5385..a00694ea5 100644 --- a/simpletuner/helpers/ramtorch/profiling.py +++ b/simpletuner/helpers/ramtorch/profiling.py @@ -80,6 +80,12 @@ def reset_for_new_run() -> None: _STATS = _initial_stats() _OUTSTANDING.clear() _STEP_DURATIONS.clear() + if torch.cuda.is_available(): + for index in range(torch.cuda.device_count()): + try: + torch.cuda.reset_peak_memory_stats(torch.device("cuda", index)) + except Exception: + continue def profile_enabled() -> bool: @@ -492,6 +498,18 @@ def snapshot() -> dict[str, Any]: } data["train_step_durations"] = _step_duration_summary() data["peak_memory"] = _peak_memory() + try: + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + get_activation_offload_pin_memory_stats, + get_activation_offload_prefetch_stats, + ) + + data["activation_offload"] = { + "pin_memory": get_activation_offload_pin_memory_stats(), + "prefetch": get_activation_offload_prefetch_stats(), + } + except Exception as exc: + data["activation_offload"] = {"error": str(exc)} data["finished_at_unix"] = time.time() return data diff --git a/simpletuner/helpers/ramtorch_extensions.py b/simpletuner/helpers/ramtorch_extensions.py index 6baad628f..b63f74464 100644 --- a/simpletuner/helpers/ramtorch_extensions.py +++ b/simpletuner/helpers/ramtorch_extensions.py @@ -23,6 +23,10 @@ _DEVICE_STATE = {} +def _forward_prefetch_stream_count() -> int: + return max(_env_int("SIMPLETUNER_RAMTORCH_FORWARD_PREFETCH_STREAMS", 4), 1) + + def _to_cpu_pinned(tensor: torch.Tensor, *, dtype: torch.dtype | None = None) -> torch.Tensor: if dtype is not None and tensor.dtype != dtype: tensor = tensor.to(dtype=dtype) @@ -47,6 +51,16 @@ def _tensor_versions(tensors): return tuple(_tensor_version(tensor) for tensor in tensors) +def _record_stream_if_supported(tensor: torch.Tensor | None, stream) -> None: + if tensor is None: + return + try: + tensor.record_stream(stream) + except (AssertionError, NotImplementedError) as exc: + if "record_stream" not in str(exc): + raise + + def _prefetch_forward_tensors(module: nn.Module, *tensors: torch.Tensor | None) -> bool: device = _device_obj(module.device) if device.type != "cuda": @@ -74,7 +88,7 @@ def _prefetch_forward_tensors(module: nn.Module, *tensors: torch.Tensor | None) return False state = _get_device_state(device) - transfer_stream = state["transfer_stream"] + transfer_stream = _next_transfer_stream(state, device) with torch.cuda.stream(transfer_stream): with record_function("forward_weight_bias_prefetch"): copied = tuple(tensor.to(device, non_blocking=True) if tensor is not None else None for tensor in tensors) @@ -110,7 +124,11 @@ def _consume_forward_prefetch(module: nn.Module, *tensors: torch.Tensor | None): return None with torch.cuda.device(device): - torch.cuda.current_stream(device).wait_event(prefetched["event"]) + current_stream = torch.cuda.current_stream(device) + current_stream.wait_event(prefetched["event"]) + for tensor in prefetched["tensors"]: + if tensor is not None: + _record_stream_if_supported(tensor, current_stream) ramtorch_profile.record_prefetch_consumed("extensions", key, device) return prefetched["tensors"] @@ -126,14 +144,18 @@ def _transfer_forward_tensors(module: nn.Module, *tensors: torch.Tensor | None): ramtorch_profile.record_fallback_forward_transfer(device, tensors) state = _get_device_state(device) - transfer_stream = state["transfer_stream"] + transfer_stream = _next_transfer_stream(state, device) with torch.cuda.stream(transfer_stream): with record_function("forward_weight_bias_transfer"): copied = tuple(tensor.to(device, non_blocking=True) if tensor is not None else None for tensor in tensors) with torch.cuda.device(device): - torch.cuda.current_stream(device).wait_stream(transfer_stream) + current_stream = torch.cuda.current_stream(device) + current_stream.wait_stream(transfer_stream) + for tensor in copied: + if tensor is not None: + _record_stream_if_supported(tensor, current_stream) return copied @@ -197,13 +219,24 @@ def _get_device_state(device): if device not in _DEVICE_STATE: with torch.cuda.device(device): _DEVICE_STATE[device] = { - "transfer_stream": torch.cuda.Stream(device=device), + "transfer_streams": [torch.cuda.Stream(device=device) for _ in range(_forward_prefetch_stream_count())], + "transfer_stream_clk": 0, "buffers": {}, "clock": 0, } return _DEVICE_STATE[device] +def _next_transfer_stream(state, device): + streams = state.get("transfer_streams") + if not streams: + streams = [state.get("transfer_stream") or torch.cuda.Stream(device=device)] + state["transfer_streams"] = streams + selected = int(state.get("transfer_stream_clk", 0)) % len(streams) + state["transfer_stream_clk"] = selected + 1 + return streams[selected] + + class CPUBouncingEmbedding(nn.Module): """ Embedding layer with CPU-stored weights that bounce to GPU on demand. @@ -1152,33 +1185,36 @@ def add_ramtorch_prefetch_hooks(module: nn.Module, component_label: str | None = hooks.append(module.register_forward_pre_hook(lambda _mod, _inp: runtime.begin_forward())) hooks.append(module.register_forward_hook(lambda _mod, _inp, _out: runtime.end_forward())) + def preserve_or_discard(current_id: str, current_module: nn.Module) -> None: + preserve_fn = getattr(current_module, "preserve_forward_for_backward", None) + if runtime.should_preserve_for_backward(current_id) and callable(preserve_fn): + preserve_fn( + max_entries=runtime.backward_preserve_max_entries, + max_bytes=runtime.backward_preserve_max_bytes, + ) + return + + discard_fn = getattr(current_module, "discard_forward_for_backward", None) + if callable(discard_fn): + discard_fn() + for module_id, current in ramtorch_modules: def pre_hook(_mod, _inp, current_id=module_id): runtime.module_entered(current_id) - return None - - def hook(_mod, _inp, _out, current_id=module_id, current_module=current): successor_id, target, source = runtime.choose_successor(current_id) if target is None or successor_id is None: ramtorch_profile.record_hook_prefetch_skipped_learned_order() - preserve_fn = getattr(current_module, "preserve_forward_for_backward", None) - if runtime.should_preserve_for_backward(current_id) and callable(preserve_fn): - preserve_fn( - max_entries=runtime.backward_preserve_max_entries, - max_bytes=runtime.backward_preserve_max_bytes, - ) return None + runtime.predicted(current_id, successor_id) ramtorch_profile.record_prefetch_successor_source(source or "traversal") success = target.prefetch_forward() ramtorch_profile.record_hook_prefetch(bool(success)) - preserve_fn = getattr(current_module, "preserve_forward_for_backward", None) - if runtime.should_preserve_for_backward(current_id) and callable(preserve_fn): - preserve_fn( - max_entries=runtime.backward_preserve_max_entries, - max_bytes=runtime.backward_preserve_max_bytes, - ) + return None + + def hook(_mod, _inp, _out, current_id=module_id, current_module=current): + preserve_or_discard(current_id, current_module) return None hooks.append(current.register_forward_pre_hook(pre_hook)) diff --git a/simpletuner/helpers/training/attention_backend.py b/simpletuner/helpers/training/attention_backend.py index 161c2d8f0..9b7615017 100644 --- a/simpletuner/helpers/training/attention_backend.py +++ b/simpletuner/helpers/training/attention_backend.py @@ -367,11 +367,30 @@ def varlen_unpacked( "flash4-hub": ("hub-fa4", "kernels-community/flash-attn4"), } +_HUB_KERNEL_VERSIONS: Dict[str, int] = { + "kernels-community/flash-attn2": 3, + "kernels-community/flash-attn3": 1, + "kernels-community/flash-attn4": 0, +} + def _normalize_backend_key(value: str) -> str: return value.replace("_", "-") +def _get_hub_kernel(target: str) -> Any: + try: + from kernels import get_kernel + except ImportError as exc: + raise RuntimeError("The 'kernels' package is required for Hugging Face Hub attention kernels.") from exc + + kwargs: Dict[str, Any] = {"trust_remote_code": True} + version = _HUB_KERNEL_VERSIONS.get(target) + if version is not None: + kwargs["version"] = version + return get_kernel(target, **kwargs) + + @lru_cache(maxsize=32) def get_packed_attention_backend( preferred_backend: Optional[str] = None, @@ -380,12 +399,7 @@ def get_packed_attention_backend( backend_key = _select_packed_backend(preferred_backend, require_varlen_qkvpacked=require_varlen_qkvpacked) provider, target = _PACKED_BACKEND_ALIASES[backend_key] if provider.startswith("hub-"): - try: - from kernels import get_kernel - except ImportError as exc: - raise RuntimeError("The 'kernels' package is required for Hugging Face Hub attention kernels.") from exc - - module = get_kernel(target) + module = _get_hub_kernel(target) return PackedAttentionBackend(backend_key, module) import importlib @@ -1120,7 +1134,8 @@ def apply(cls, config, phase: AttentionPhase) -> None: diffusers_backend = cls._resolve_diffusers_backend(backend_alias) if diffusers_backend is not None: cls._clear_metal_flash_attention_quantization_mode() - cls._enable_diffusers_backend(backend_alias, diffusers_backend) + trust_remote_code = cls._truthy_config_value(getattr(config, "trust_remote_code", False)) + cls._enable_diffusers_backend(backend_alias, diffusers_backend, trust_remote_code=trust_remote_code) return cls.restore_default() @@ -1168,6 +1183,14 @@ def _resolve_diffusers_backend(cls, backend: str) -> Optional[AttentionBackendNa normalized = cls._normalize_backend_key(backend) return _DIFFUSERS_BACKEND_ALIASES.get(normalized) + @staticmethod + def _truthy_config_value(value: Any) -> bool: + if isinstance(value, bool): + return value + if isinstance(value, str): + return value.strip().lower() in {"1", "true", "yes", "on"} + return bool(value) + @classmethod def _load_metal_flash_attention_extension(cls, backend: str = "metal-flash-attention"): backend_key = _normalize_backend_key(str(backend or "metal-flash-attention").strip().lower()) @@ -1358,29 +1381,70 @@ def _metal_flash_attention_should_fallback(query, key, value, attn_mask, dropout return False @classmethod - def _enable_diffusers_backend(cls, backend_key: str, backend_enum: AttentionBackendName) -> None: + def _enable_diffusers_backend( + cls, + backend_key: str, + backend_enum: AttentionBackendName, + *, + trust_remote_code: bool = False, + ) -> None: if cls._diffusers_backend_name == backend_key: return if diffusers_attention_backend is None or _check_attention_backend_requirements is None: - message = f"Diffusers attention backend helpers are unavailable. Upgrade diffusers to at least 0.35 to use {backend_key}." + message = ( + "Diffusers attention backend helpers are unavailable. " + f"Upgrade diffusers to at least 0.35 to use {backend_key}." + ) logger.error(message) raise RuntimeError(message) - try: - _check_attention_backend_requirements(backend_enum) - except Exception as exc: # pragma: no cover - exercised only when backend requirements fail - message = f"Attention backend '{backend_key}' is unavailable: {exc}" - logger.error(message) - raise RuntimeError(message) from exc + patched_kernel_modules: list[tuple[Any, Any]] = [] + if trust_remote_code and backend_key.endswith("-hub"): + for module_name in ("kernels", "kernels.utils"): + try: + kernel_module = importlib.import_module(module_name) + original_get_kernel = getattr(kernel_module, "get_kernel", None) + except ModuleNotFoundError as exc: + missing_module = exc.name or "" + if missing_module not in {module_name, module_name.split(".", 1)[0]}: + raise + continue + if not callable(original_get_kernel): + continue + + def patched_get_kernel(*args, _original_get_kernel=original_get_kernel, **kwargs): + kwargs.setdefault("trust_remote_code", True) + repo_id = kwargs.get("repo_id") + if repo_id is None and args: + repo_id = args[0] + if kwargs.get("version") is None and kwargs.get("revision") is None: + version = _HUB_KERNEL_VERSIONS.get(repo_id) + if version is not None: + kwargs["version"] = version + return _original_get_kernel(*args, **kwargs) + + setattr(kernel_module, "get_kernel", patched_get_kernel) + patched_kernel_modules.append((kernel_module, original_get_kernel)) - cls._disable_diffusers_backend() try: - context = diffusers_attention_backend(backend_enum) - context.__enter__() - except Exception as exc: - message = f"Failed to enable attention backend '{backend_key}': {exc}" - logger.error(message) - raise RuntimeError(message) from exc + try: + _check_attention_backend_requirements(backend_enum) + except Exception as exc: # pragma: no cover - exercised only when backend requirements fail + message = f"Attention backend '{backend_key}' is unavailable: {exc}" + logger.error(message) + raise RuntimeError(message) from exc + + cls._disable_diffusers_backend() + try: + context = diffusers_attention_backend(backend_enum) + context.__enter__() + except Exception as exc: + message = f"Failed to enable attention backend '{backend_key}': {exc}" + logger.error(message) + raise RuntimeError(message) from exc + finally: + for kernel_module, original_get_kernel in patched_kernel_modules: + setattr(kernel_module, "get_kernel", original_get_kernel) cls._diffusers_backend_context = context cls._diffusers_backend_name = backend_key @@ -1785,7 +1849,8 @@ def _load_state_store(cls, payload: Dict[str, Any]) -> None: saved_settings = payload.get("settings") if saved_settings and cls._sla_settings and cls._sla_settings != saved_settings: logger.warning( - "SLA runtime settings differ from checkpoint settings. Runtime=%s, Checkpoint=%s. Proceeding with runtime configuration.", + "SLA runtime settings differ from checkpoint settings. Runtime=%s, Checkpoint=%s. " + "Proceeding with runtime configuration.", cls._sla_settings, saved_settings, ) diff --git a/simpletuner/helpers/training/collate.py b/simpletuner/helpers/training/collate.py index 819eec7fe..b13f0600f 100644 --- a/simpletuner/helpers/training/collate.py +++ b/simpletuner/helpers/training/collate.py @@ -396,6 +396,7 @@ def _collate_tensors(tensors): - If tensors are 3D [1, seq, dim], concatenate along dim=0 to get [batch, seq, dim] - If tensors have inconsistent dimensions, normalize them first """ + tensors = [tensor for tensor in tensors if tensor is not None] if not tensors: return None @@ -645,6 +646,13 @@ def collate_fn(batch): token_payload["is_i2v_data"] = is_i2v_data return token_payload + uses_text_embeddings_cache = True + if model is not None: + try: + uses_text_embeddings_cache = bool(model.uses_text_embeddings_cache()) + except AttributeError: + uses_text_embeddings_cache = True + debug_log("Compute latents") batch_data = compute_latents(filepaths, batch_backend_id, model) latent_metadata = None @@ -1026,17 +1034,16 @@ def _conditioning_pixel_value_for_example(example_idx: int): # If the caption is empty, we use the instance prompt text. captions = [caption if caption else example["instance_prompt_text"] for caption, example in zip(captions, examples)] debug_log(f"Pull cached text embeds. conditioning captions: {captions}") - - # Get the appropriate text_embed_cache - if conditioning_backends: - text_embed_cache = conditioning_backends[0]["text_embed_cache"] - else: - text_embed_cache = StateTracker.get_data_backend(data_backend_id)["text_embed_cache"] else: # Use training captions (default behavior) captions = [example["instance_prompt_text"] for example in examples] debug_log(f"Pull cached text embeds. Using training set captions: {captions}") - text_embed_cache = StateTracker.get_data_backend(data_backend_id)["text_embed_cache"] + text_embed_cache = None + if uses_text_embeddings_cache: + if has_conditioning_captions and is_random_mode and conditioning_backends: + text_embed_cache = conditioning_backends[0]["text_embed_cache"] + else: + text_embed_cache = StateTracker.get_data_backend(data_backend_id)["text_embed_cache"] prompt_requests = [] key_type = TextEmbedCacheKey.CAPTION getter = getattr(model, "text_embed_cache_key", None) @@ -1046,53 +1053,54 @@ def _conditioning_pixel_value_for_example(example_idx: int): except Exception as exc: debug_log(f"text_embed_cache_key() lookup failed on model {type(model)}: {exc}") - for idx, caption in enumerate(captions): - example = examples[idx] - example_path = example.get("image_path") - example_backend_id = example.get("data_backend_id") - backend_config = StateTracker.get_data_backend_config(example_backend_id) if example_backend_id else {} - backend_config = backend_config or {} - dataset_root = backend_config.get("instance_data_dir") - normalized_identifier = normalize_data_path(example_path, dataset_root) - metadata = { - "image_path": example_path, - "data_backend_id": example_backend_id, - "prompt": caption, - "dataset_relative_path": normalized_identifier, - } - metadata_builder = getattr(model, "text_embed_cache_metadata_for_sample", None) - if callable(metadata_builder): - metadata.update( - metadata_builder( - example=example, - latent=latent_batch[idx], - prompt=caption, - data_backend_id=example_backend_id, - dataset_relative_path=normalized_identifier, + if uses_text_embeddings_cache: + for idx, caption in enumerate(captions): + example = examples[idx] + example_path = example.get("image_path") + example_backend_id = example.get("data_backend_id") + backend_config = StateTracker.get_data_backend_config(example_backend_id) if example_backend_id else {} + backend_config = backend_config or {} + dataset_root = backend_config.get("instance_data_dir") + normalized_identifier = normalize_data_path(example_path, dataset_root) + metadata = { + "image_path": example_path, + "data_backend_id": example_backend_id, + "prompt": caption, + "dataset_relative_path": normalized_identifier, + } + metadata_builder = getattr(model, "text_embed_cache_metadata_for_sample", None) + if callable(metadata_builder): + metadata.update( + metadata_builder( + example=example, + latent=latent_batch[idx], + prompt=caption, + data_backend_id=example_backend_id, + dataset_relative_path=normalized_identifier, + ) ) - ) - # Only include conditioning pixels for text embedding when using a single - # conditioning image. With multiple backends in combined mode, skip image - # context in embeddings and rely solely on latent references. - # (In random mode with multiple backends, only one image is selected, so - # we can still use it for text embedding context.) - has_multiple_combined_refs = len(conditioning_backends) > 1 and is_combined_mode - if not has_multiple_combined_refs: - pixel_value = _conditioning_pixel_value_for_example(idx) - if pixel_value is not None: - metadata["conditioning_pixel_values"] = pixel_value - if key_type is TextEmbedCacheKey.DATASET_AND_FILENAME and example_backend_id and example_path: - key_value = f"{example_backend_id}:{normalized_identifier}" - elif key_type is TextEmbedCacheKey.FILENAME and example_path: - key_value = normalize_data_path(example_path, None) - else: - key_value = caption - key_builder = getattr(model, "text_embed_cache_key_value", None) - if callable(key_builder): - key_value = key_builder(prompt=caption, default_key=key_value, metadata=metadata) - prompt_requests.append({"prompt": caption, "key": key_value, "metadata": metadata}) + # Only include conditioning pixels for text embedding when using a single + # conditioning image. With multiple backends in combined mode, skip image + # context in embeddings and rely solely on latent references. + # (In random mode with multiple backends, only one image is selected, so + # we can still use it for text embedding context.) + has_multiple_combined_refs = len(conditioning_backends) > 1 and is_combined_mode + if not has_multiple_combined_refs: + pixel_value = _conditioning_pixel_value_for_example(idx) + if pixel_value is not None: + metadata["conditioning_pixel_values"] = pixel_value + if key_type is TextEmbedCacheKey.DATASET_AND_FILENAME and example_backend_id and example_path: + key_value = f"{example_backend_id}:{normalized_identifier}" + elif key_type is TextEmbedCacheKey.FILENAME and example_path: + key_value = normalize_data_path(example_path, None) + else: + key_value = caption + key_builder = getattr(model, "text_embed_cache_key_value", None) + if callable(key_builder): + key_value = key_builder(prompt=caption, default_key=key_value, metadata=metadata) + prompt_requests.append({"prompt": caption, "key": key_value, "metadata": metadata}) - if not text_embed_cache.disabled: + if text_embed_cache is not None and not text_embed_cache.disabled: all_text_encoder_outputs = compute_prompt_embeddings(prompt_requests, text_embed_cache, StateTracker.get_model()) else: all_text_encoder_outputs = {} @@ -1132,6 +1140,19 @@ def _conditioning_pixel_value_for_example(example_idx: int): if not any(s2v_audio_paths): backend_config = StateTracker.get_data_backend_config(batch_backend_id) or {} dataset_type = backend_config.get("dataset_type") + + def resolve_local_source_audio_path(path: str | None) -> str | None: + if not path: + return None + candidates = [path] + instance_data_dir = backend_config.get("instance_data_dir") + if instance_data_dir and not os.path.isabs(path): + candidates.insert(0, os.path.join(instance_data_dir, path.lstrip(os.sep))) + for candidate in candidates: + if os.path.exists(candidate): + return candidate + return None + if dataset_type == "video": s2v_datasets = StateTracker.get_s2v_datasets(batch_backend_id) if len(s2v_datasets) == 1: @@ -1142,17 +1163,20 @@ def _conditioning_pixel_value_for_example(example_idx: int): if audio_backend_id is not None: metadata_backend = s2v_datasets[0].get("metadata_backend") if metadata_backend is None: - s2v_audio_paths = list(filepaths) - s2v_audio_backend_ids = [audio_backend_id] * len(s2v_audio_paths) + s2v_audio_paths = [resolve_local_source_audio_path(path) for path in filepaths] + s2v_audio_backend_ids = [ + audio_backend_id if path is not None else None for path in s2v_audio_paths + ] else: s2v_audio_paths = [] s2v_audio_backend_ids = [] for path in filepaths: - if metadata_backend.get_metadata_by_filepath(path) is None: + audio_path = resolve_local_source_audio_path(path) + if metadata_backend.get_metadata_by_filepath(path) is None or audio_path is None: s2v_audio_paths.append(None) s2v_audio_backend_ids.append(None) else: - s2v_audio_paths.append(path) + s2v_audio_paths.append(audio_path) s2v_audio_backend_ids.append(audio_backend_id) audio_latent_batch = None @@ -1228,6 +1252,8 @@ def _conditioning_pixel_value_for_example(example_idx: int): grounding_batch = None max_grounding_entities = getattr(StateTracker.get_args(), "max_grounding_entities", None) if max_grounding_entities and max_grounding_entities > 0: + if text_embed_cache is None: + raise ValueError("Grounding annotations require a text embedding cache.") from simpletuner.helpers.training.grounding.collate import GroundingCollate grounding_image_cache = StateTracker.get_grounding_image_embed_cache(data_backend_id) diff --git a/simpletuner/helpers/training/default_settings/safety_check.py b/simpletuner/helpers/training/default_settings/safety_check.py index d587f92bb..da0ba38dc 100644 --- a/simpletuner/helpers/training/default_settings/safety_check.py +++ b/simpletuner/helpers/training/default_settings/safety_check.py @@ -8,6 +8,7 @@ from simpletuner.helpers.training.attention_backend import AttentionBackendMode from simpletuner.helpers.training.multi_process import _get_rank as get_rank +from simpletuner.helpers.training.offloaded_gradient_checkpointer import normalize_activation_offload_pin_memory_max_buckets logger = logging.getLogger(__name__) from simpletuner.helpers.training.multi_process import should_log @@ -141,6 +142,7 @@ def safety_check(args, accelerator): gradient_checkpointing_interval_supported_models = [ "flux", + "ace_step", "sana", "sd3", "chroma", @@ -148,7 +150,94 @@ def safety_check(args, accelerator): "hunyuanvideo", "ernie", "cosmos3", + "mageflow", + "boogu_image", + "cosmos", + "flux2", + "hidream", + "ideogram", + "kandinsky5_image", + "kandinsky5_video", + "krea2", + "longcat_image", + "longcat_video", + "ltxvideo", + "ltxvideo2", + "lumina2", + "pixart", + "qwen_image", + "sanavideo", + "stable_cascade", + "wan_s2v", + "wan", + "z_image", + "zlab_i1", ] + gradient_checkpointing_segment_stride_supported_models = [ + "ace_step", + "auraflow", + "boogu_image", + "chroma", + "cosmos", + "cosmos3", + "ernie", + "flux", + "flux2", + "hidream", + "hunyuanvideo", + "ideogram", + "kandinsky5_image", + "kandinsky5_video", + "krea2", + "longcat_image", + "longcat_video", + "ltxvideo", + "ltxvideo2", + "lumina2", + "mageflow", + "pixart", + "qwen_image", + "sana", + "sanavideo", + "sd3", + "stable_cascade", + "wan_s2v", + "wan", + "z_image", + "zlab_i1", + ] + attention_activation_offload_supported_models = [ + "chroma", + "flux", + "flux2", + "hunyuanvideo", + "kandinsky5-image", + "kandinsky5-video", + "kandinsky5_image", + "kandinsky5_video", + "krea2", + "longcat_image", + "longcat_video", + "ltxvideo2", + "mageflow", + "sd3", + "wan", + "z_image", + ] + if getattr(args, "gradient_checkpointing_offload_attention", False): + if args.model_family.lower() not in attention_activation_offload_supported_models: + raise ValueError( + "--gradient_checkpointing_offload_attention is only supported on model families with a clean " + f"attention/FFN checkpointing boundary. Currently supported models: {attention_activation_offload_supported_models}" + ) + elif getattr(args, "gradient_checkpointing_offload_prefetch", False): + logger.warning( + "Gradient checkpointing activation prefetch requires --gradient_checkpointing_offload_attention; disabling prefetch." + ) + args.gradient_checkpointing_offload_prefetch = False + args.gradient_checkpointing_offload_pin_memory_max_buckets = normalize_activation_offload_pin_memory_max_buckets( + getattr(args, "gradient_checkpointing_offload_pin_memory_max_buckets", 12) + ) if args.gradient_checkpointing_interval == 1: args.gradient_checkpointing_interval = None if args.gradient_checkpointing_interval is not None: @@ -159,6 +248,28 @@ def safety_check(args, accelerator): args.gradient_checkpointing_interval = None if args.gradient_checkpointing_interval == 0: raise ValueError("Gradient checkpointing interval must be greater than 0. Please set it to a positive integer.") + if getattr(args, "gradient_checkpointing_segment_stride", None) in ("", "None"): + args.gradient_checkpointing_segment_stride = None + if getattr(args, "gradient_checkpointing_segment_stride", None) is not None: + try: + args.gradient_checkpointing_segment_stride = int(args.gradient_checkpointing_segment_stride) + except (TypeError, ValueError): + raise ValueError("Gradient checkpointing segment stride must be a positive integer.") + if args.model_family.lower() not in gradient_checkpointing_segment_stride_supported_models: + logger.warning( + f"Gradient checkpointing segment stride is not supported with {args.model_family} models. " + f"Currently supported models: {gradient_checkpointing_segment_stride_supported_models}" + ) + args.gradient_checkpointing_segment_stride = None + elif args.gradient_checkpointing_segment_stride <= 0: + raise ValueError("Gradient checkpointing segment stride must be greater than 0.") + elif args.gradient_checkpointing_interval is None: + logger.warning( + "Gradient checkpointing segment stride requires --gradient_checkpointing_interval greater than 1; ignoring segment stride." + ) + args.gradient_checkpointing_segment_stride = None + elif args.gradient_checkpointing_segment_stride < args.gradient_checkpointing_interval: + raise ValueError("Gradient checkpointing segment stride must be at least the checkpointing interval.") def _normalize_interval(raw_value, cast): if raw_value in (None, "", "None"): diff --git a/simpletuner/helpers/training/gradient_checkpointing_interval.py b/simpletuner/helpers/training/gradient_checkpointing_interval.py index ec1c42bb1..b468a6d4a 100644 --- a/simpletuner/helpers/training/gradient_checkpointing_interval.py +++ b/simpletuner/helpers/training/gradient_checkpointing_interval.py @@ -1,13 +1,7 @@ -""" -Gradient checkpointing backend selection. +"""Gradient checkpointing backend and segmentation helpers.""" -This module provides the ability to select between different gradient checkpointing -backends (torch native vs unsloth CPU offload). - -Note: Per-layer interval checkpointing is implemented directly in transformer models -that support it (Flux, Chroma, SD3, Sana, AuraFlow, etc.) via their -`set_gradient_checkpointing_interval` method. -""" +from collections.abc import Callable, Sequence +from typing import Any _checkpoint_backend = "torch" # "torch", "unsloth", "torch-ffn", or "unsloth-ffn" _offloaded_checkpoint = None # Lazy import @@ -53,3 +47,74 @@ def get_checkpoint_function(): if get_checkpoint_backend_base() == "unsloth" and _offloaded_checkpoint is not None: return _offloaded_checkpoint return torch.utils.checkpoint.checkpoint + + +def should_checkpoint_block( + block_index: int, + gradient_checkpointing: bool, + interval: int | None = None, + segment_stride: int | None = None, +) -> bool: + """Return whether an individual block should be checkpointed. + + ``interval=None`` means every block, matching the legacy layer checkpointing + behavior. With a stride, checkpoint the first ``interval`` blocks in each + stride window, e.g. interval=2 stride=4 checkpoints 0,1,4,5,... + """ + if not gradient_checkpointing: + return False + if interval is None or interval <= 1: + return True + if segment_stride is None: + return block_index % interval == 0 + if segment_stride < interval: + raise ValueError("segment_stride must be at least interval") + return block_index % segment_stride < interval + + +def checkpoint_sequential_state( + blocks: Sequence[Any], + segment_size: int, + state: tuple[Any, ...] | Any, + run_block: Callable[..., tuple[Any, ...] | Any], + checkpoint_fn: Callable[..., Any], + checkpoint_kwargs: dict[str, Any] | None = None, + segment_stride: int | None = None, +) -> tuple[Any, ...]: + """Checkpoint contiguous chunks of a stateful block sequence. + + Unlike ``torch.utils.checkpoint.checkpoint_sequential``, this supports block + functions that carry multiple tensors through the sequence. + """ + if segment_size < 1: + raise ValueError("segment_size must be greater than 0") + if segment_stride is None: + segment_stride = segment_size + if segment_stride < segment_size: + raise ValueError("segment_stride must be at least segment_size") + + current_state = state if isinstance(state, tuple) else (state,) + checkpoint_kwargs = dict(checkpoint_kwargs or {}) + + for segment_start in range(0, len(blocks), segment_stride): + segment_blocks = tuple(blocks[segment_start : segment_start + segment_size]) + + def run_segment(*segment_state, _segment_start=segment_start, _segment_blocks=segment_blocks): + next_state = segment_state + for offset, block in enumerate(_segment_blocks): + result = run_block(_segment_start + offset, block, *next_state) + next_state = result if isinstance(result, tuple) else (result,) + if len(next_state) == 1: + return next_state[0] + return next_state + + result = checkpoint_fn(run_segment, *current_state, **checkpoint_kwargs) + current_state = result if isinstance(result, tuple) else (result,) + + gap_start = segment_start + len(segment_blocks) + gap_end = min(segment_start + segment_stride, len(blocks)) + for block_index in range(gap_start, gap_end): + result = run_block(block_index, blocks[block_index], *current_state) + current_state = result if isinstance(result, tuple) else (result,) + + return current_state diff --git a/simpletuner/helpers/training/lora_format.py b/simpletuner/helpers/training/lora_format.py index 01a2954a1..28e4c588e 100644 --- a/simpletuner/helpers/training/lora_format.py +++ b/simpletuner/helpers/training/lora_format.py @@ -34,13 +34,10 @@ def detect_state_dict_format(state_dict: Dict[str, Any]) -> Optional[PEFTLoRAFor keys = list(state_dict.keys()) comfy_prefix_hits = sum(k.startswith("diffusion_model.") for k in keys) comfy_alpha_hits = sum(k.endswith(".alpha") for k in keys) - comfy_ab_hits = sum(".lora_A" in k or ".lora_B" in k for k in keys) diffusers_down_up_hits = sum(".lora.down" in k or ".lora.up" in k for k in keys) if comfy_prefix_hits or (comfy_alpha_hits and diffusers_down_up_hits == 0): return PEFTLoRAFormat.COMFYUI - if comfy_ab_hits and diffusers_down_up_hits == 0 and comfy_prefix_hits >= 0: - return PEFTLoRAFormat.COMFYUI return PEFTLoRAFormat.DIFFUSERS @@ -233,24 +230,16 @@ def convert_diffusers_to_comfyui( if ".lora.down." in new_key: new_key = new_key.replace(".lora.down.", ".lora_A.") + elif ".lora.up." in new_key: + new_key = new_key.replace(".lora.up.", ".lora_B.") + + if ".lora_A." in new_key: module_key = new_key[: new_key.rfind(".lora_A.")] alpha_value = _resolve_alpha_for_module( module_key.removeprefix(f"{diffusion_prefix}."), weight, adapter_metadata ) if alpha_value is not None and module_key not in alpha_entries: alpha_entries[module_key] = torch.tensor(alpha_value, dtype=torch.float32) - elif new_key.endswith(".lora.down.weight"): - new_key = new_key.replace(".lora.down.weight", ".lora_A.weight") - module_key = new_key[: new_key.rfind(".lora_A.weight")] - alpha_value = _resolve_alpha_for_module( - module_key.removeprefix(f"{diffusion_prefix}."), weight, adapter_metadata - ) - if alpha_value is not None and module_key not in alpha_entries: - alpha_entries[module_key] = torch.tensor(alpha_value, dtype=torch.float32) - elif ".lora.up." in new_key: - new_key = new_key.replace(".lora.up.", ".lora_B.") - elif new_key.endswith(".lora.up.weight"): - new_key = new_key.replace(".lora.up.weight", ".lora_B.weight") converted[new_key] = weight diff --git a/simpletuner/helpers/training/offloaded_gradient_checkpointer.py b/simpletuner/helpers/training/offloaded_gradient_checkpointer.py index 2b7b65cc9..aab525323 100644 --- a/simpletuner/helpers/training/offloaded_gradient_checkpointer.py +++ b/simpletuner/helpers/training/offloaded_gradient_checkpointer.py @@ -1,22 +1,740 @@ -""" -Unsloth-style gradient checkpointing using CPU offload. +"""CPU saved-tensor offload helpers for activation memory pressure.""" -Uses PyTorch's saved_tensors_hooks to intercept tensor saves during checkpoint, -offloading them to CPU asynchronously and restoring during backward pass. -This trades PCIe bandwidth for GPU memory savings. -""" +import sys +from collections import OrderedDict, defaultdict +from contextlib import nullcontext +from dataclasses import dataclass import torch from diffusers.utils.torch_utils import is_torch_version from torch.utils.checkpoint import checkpoint as torch_checkpoint +_DEFAULT_PIN_MEMORY_MAX_BUCKETS = 12 +_ACTIVATION_PREFETCH_MIN_OBSERVATIONS = 2 +_ACTIVATION_PREFETCH_MIN_CONFIDENCE = 0.80 +_ACTIVATION_PREFETCH_AUTOTUNE_MIN_SAMPLES = 8 +_ACTIVATION_PREFETCH_AUTOTUNE_MARGIN = 0.98 +_DEFAULT_D2H_COPY_STREAMS = 2 +_DEFAULT_H2D_PREFETCH_STREAMS = 4 + + +@dataclass(frozen=True) +class _PinnedBucketKey: + size: tuple[int, ...] + stride: tuple[int, ...] + dtype: torch.dtype + layout: torch.layout + + +@dataclass(frozen=True) +class _RestoreView: + size: tuple[int, ...] + stride: tuple[int, ...] + storage_offset: int + + +@dataclass +class _OffloadedActivationRecord: + logical_id: str + predictor_id: str + generation: int + tensor: torch.Tensor + original_device: torch.device + pool_key: "_PinnedBucketKey | None" + restore_view: _RestoreView | None + ready_event: torch.cuda.Event | None + prefetched_tensor: torch.Tensor | None = None + prefetch_event: torch.cuda.Event | None = None + consumed: bool = False + cpu_released: bool = False + + +@dataclass +class _PinnedBucketStats: + accesses: int = 0 + resident_accesses: int = 0 + pinned_checkouts: int = 0 + buffer_reuses: int = 0 + allocations: int = 0 + admissions: int = 0 + evictions: int = 0 + cap_misses: int = 0 + disabled_misses: int = 0 + allocation_failures: int = 0 + releases: int = 0 + dropped_releases: int = 0 + last_access: int = 0 + last_admission: int = 0 + last_eviction: int = 0 + + @property + def pinned_checkout_rate(self) -> float: + return self.pinned_checkouts / self.accesses if self.accesses else 0.0 + + @property + def buffer_reuse_rate(self) -> float: + return self.buffer_reuses / self.pinned_checkouts if self.pinned_checkouts else 0.0 + + +class _PinnedMemoryPool: + def __init__(self, max_buckets: int = _DEFAULT_PIN_MEMORY_MAX_BUCKETS): + self.max_buckets = max(0, int(max_buckets)) + self.available: OrderedDict[_PinnedBucketKey, list[torch.Tensor]] = OrderedDict() + self.pending: dict[_PinnedBucketKey, list[tuple[torch.cuda.Event, torch.Tensor]]] = defaultdict(list) + self.stats: dict[_PinnedBucketKey, _PinnedBucketStats] = {} + self.total_accesses = 0 + self.total_pinned_checkouts = 0 + self.total_buffer_reuses = 0 + self.total_allocations = 0 + self.total_cap_misses = 0 + self.total_evictions = 0 + + def set_max_buckets(self, max_buckets: int) -> None: + self.max_buckets = max(0, int(max_buckets)) + while len(self.available) > self.max_buckets: + victim, _ = next(iter(self.available.items())) + self._evict(victim) + + def key_for(self, tensor: torch.Tensor) -> _PinnedBucketKey | None: + if tensor.layout != torch.strided: + return None + return _PinnedBucketKey(tuple(tensor.size()), tuple(tensor.stride()), tensor.dtype, tensor.layout) + + def checkout(self, key: _PinnedBucketKey) -> torch.Tensor | None: + stats = self._record_access(key) + if self.max_buckets <= 0: + stats.disabled_misses += 1 + return None + self._drain_completed(key) + if key not in self.available: + if len(self.available) >= self.max_buckets: + if not self._admit_by_eviction(key): + stats.cap_misses += 1 + self.total_cap_misses += 1 + return None + self.available[key] = [] + stats.admissions += 1 + stats.last_admission = self.total_accesses + else: + stats.resident_accesses += 1 + self.available.move_to_end(key) + if self.available[key]: + stats.buffer_reuses += 1 + stats.pinned_checkouts += 1 + self.total_buffer_reuses += 1 + self.total_pinned_checkouts += 1 + return self.available[key].pop() + try: + tensor = self._allocate(key) + except RuntimeError: + stats.allocation_failures += 1 + self.discard_empty_bucket(key) + return None + stats.allocations += 1 + stats.pinned_checkouts += 1 + self.total_allocations += 1 + self.total_pinned_checkouts += 1 + return tensor + + def release_after_cuda_copy(self, key: _PinnedBucketKey, tensor: torch.Tensor, device: torch.device) -> None: + stats = self._stats_for(key) + if key not in self.available: + stats.dropped_releases += 1 + return + stats.releases += 1 + if device.type == "cuda" and torch.cuda.is_available(): + event = torch.cuda.Event() + event.record(torch.cuda.current_stream(device)) + self._release_after_event(key, tensor, event) + return + self.available[key].append(tensor) + + def release_after_record_event( + self, key: _PinnedBucketKey, tensor: torch.Tensor, event: torch.cuda.Event | None + ) -> None: + stats = self._stats_for(key) + if key not in self.available: + stats.dropped_releases += 1 + return + stats.releases += 1 + self._release_after_event(key, tensor, event) + + def _release_after_event(self, key: _PinnedBucketKey, tensor: torch.Tensor, event: torch.cuda.Event | None) -> None: + if event is not None and torch.cuda.is_available(): + self.pending[key].append((event, tensor)) + return + self.available[key].append(tensor) + + def discard_empty_bucket(self, key: _PinnedBucketKey) -> None: + if not self.available.get(key) and not self.pending.get(key): + self.available.pop(key, None) + + def _drain_completed(self, key: _PinnedBucketKey) -> None: + pending = self.pending.get(key) + if not pending: + return + remaining = [] + for event, tensor in pending: + if event.query(): + self.available.setdefault(key, []).append(tensor) + else: + remaining.append((event, tensor)) + if remaining: + self.pending[key] = remaining + else: + self.pending.pop(key, None) + + def _allocate(self, key: _PinnedBucketKey) -> torch.Tensor: + return torch.empty_strided(key.size, key.stride, dtype=key.dtype, layout=key.layout, device="cpu", pin_memory=True) + + def _record_access(self, key: _PinnedBucketKey) -> _PinnedBucketStats: + stats = self._stats_for(key) + self.total_accesses += 1 + stats.accesses += 1 + stats.last_access = self.total_accesses + return stats + + def _stats_for(self, key: _PinnedBucketKey) -> _PinnedBucketStats: + stats = self.stats.get(key) + if stats is None: + stats = _PinnedBucketStats() + self.stats[key] = stats + return stats + + def _admit_by_eviction(self, candidate: _PinnedBucketKey) -> bool: + evictable = [key for key in self.available if not self.pending.get(key)] + if not evictable: + return False + candidate_score = self._admission_score(candidate) + victim = min(evictable, key=self._admission_score) + if candidate_score <= self._admission_score(victim): + return False + self._evict(victim) + return True + + def _admission_score(self, key: _PinnedBucketKey) -> tuple[int, float, int]: + stats = self._stats_for(key) + return (stats.accesses, stats.pinned_checkout_rate, stats.last_access) + + def _evict(self, key: _PinnedBucketKey) -> None: + self.available.pop(key, None) + stats = self._stats_for(key) + stats.evictions += 1 + stats.last_eviction = self.total_accesses + self.total_evictions += 1 + + def snapshot(self) -> dict: + buckets = [] + resident = set(self.available) + pending_counts = {key: len(value) for key, value in self.pending.items()} + for key, stats in self.stats.items(): + buckets.append( + { + "size": key.size, + "stride": key.stride, + "dtype": str(key.dtype), + "layout": str(key.layout), + "resident": key in resident, + "available_buffers": len(self.available.get(key, ())), + "pending_buffers": pending_counts.get(key, 0), + "accesses": stats.accesses, + "resident_accesses": stats.resident_accesses, + "pinned_checkouts": stats.pinned_checkouts, + "buffer_reuses": stats.buffer_reuses, + "allocations": stats.allocations, + "admissions": stats.admissions, + "evictions": stats.evictions, + "cap_misses": stats.cap_misses, + "disabled_misses": stats.disabled_misses, + "allocation_failures": stats.allocation_failures, + "releases": stats.releases, + "dropped_releases": stats.dropped_releases, + "pinned_checkout_rate": stats.pinned_checkout_rate, + "buffer_reuse_rate": stats.buffer_reuse_rate, + "last_access": stats.last_access, + "last_admission": stats.last_admission, + "last_eviction": stats.last_eviction, + } + ) + buckets.sort(key=lambda item: (item["accesses"], item["pinned_checkouts"], item["last_access"]), reverse=True) + return { + "max_buckets": self.max_buckets, + "resident_buckets": len(self.available), + "tracked_buckets": len(self.stats), + "total_accesses": self.total_accesses, + "total_pinned_checkouts": self.total_pinned_checkouts, + "total_buffer_reuses": self.total_buffer_reuses, + "total_allocations": self.total_allocations, + "total_cap_misses": self.total_cap_misses, + "total_evictions": self.total_evictions, + "pinned_checkout_rate": self.total_pinned_checkouts / self.total_accesses if self.total_accesses else 0.0, + "buffer_reuse_rate": ( + self.total_buffer_reuses / self.total_pinned_checkouts if self.total_pinned_checkouts else 0.0 + ), + "buckets": buckets, + } + + def reset_stats(self) -> None: + self.stats.clear() + self.total_accesses = 0 + self.total_pinned_checkouts = 0 + self.total_buffer_reuses = 0 + self.total_allocations = 0 + self.total_cap_misses = 0 + self.total_evictions = 0 + + +_PINNED_MEMORY_POOL = _PinnedMemoryPool() +_ACTIVATION_PREFETCH_ENABLED = False +_ACTIVATION_PREFETCH_AUTOTUNE_ENABLED = True + + +class _CudaStreamPool: + def __init__(self, width: int): + self.width = max(1, int(width)) + self.streams: dict[int, list[torch.cuda.Stream]] = {} + self.next_index: dict[int, int] = defaultdict(int) + self.uses: dict[int, list[int]] = {} + + def set_width(self, width: int) -> None: + self.width = max(1, int(width)) + self.streams.clear() + self.next_index.clear() + self.uses.clear() + + def next(self, device: torch.device) -> torch.cuda.Stream: + device_index = _device_index(device) + streams = self.streams.get(device_index) + if streams is None or len(streams) != self.width: + streams = [torch.cuda.Stream(device=device_index) for _ in range(self.width)] + self.streams[device_index] = streams + self.next_index[device_index] = 0 + self.uses[device_index] = [0 for _ in range(self.width)] + stream_index = self.next_index[device_index] + self.next_index[device_index] = (stream_index + 1) % self.width + self.uses[device_index][stream_index] += 1 + return streams[stream_index] + + def snapshot(self) -> dict: + return { + "width": self.width, + "devices": { + device_index: { + "allocated_streams": len(self.streams.get(device_index, ())), + "uses": list(self.uses.get(device_index, ())), + "total_uses": sum(self.uses.get(device_index, ())), + } + for device_index in sorted(set(self.streams) | set(self.uses)) + }, + } + + def reset_stats(self) -> None: + for device_index, uses in list(self.uses.items()): + self.uses[device_index] = [0 for _ in uses] + + +_D2H_COPY_STREAMS = _CudaStreamPool(_DEFAULT_D2H_COPY_STREAMS) +_H2D_PREFETCH_STREAMS = _CudaStreamPool(_DEFAULT_H2D_PREFETCH_STREAMS) + + +class _ActivationOffloadPrefetchRuntime: + def __init__(self): + self.records: dict[tuple[int, str], list[_OffloadedActivationRecord]] = defaultdict(list) + self.transitions: dict[str, dict[str, int]] = defaultdict(dict) + self.successors: dict[str, str] = {} + self.disabled: dict[str, dict] = {} + self.previous_by_generation: dict[int, str] = {} + self.current_generation = 0 + self.saw_unpack_since_pack = False + self.total_packs = 0 + self.total_unpacks = 0 + self.prefetch_attempts = 0 + self.prefetch_hits = 0 + self.prefetch_misses = 0 + self.prefetch_enqueued = 0 + self.prefetch_skipped = 0 + self.prefetch_stale = 0 + self.transition_updates = 0 + self.jit_restore_ms: list[float] = [] + self.prefetch_wait_ms: list[float] = [] + self.autotune_disabled = False + self.autotune_decision: str | None = None + + def next_generation_for_pack(self) -> int: + if self.saw_unpack_since_pack: + self._retire_generation(self.current_generation) + self.current_generation += 1 + self.saw_unpack_since_pack = False + self.previous_by_generation.pop(self.current_generation, None) + return self.current_generation + + def register(self, record: _OffloadedActivationRecord) -> None: + self.total_packs += 1 + self.records[(record.generation, record.predictor_id)].append(record) + + def consume(self, record: _OffloadedActivationRecord) -> None: + self.total_unpacks += 1 + self.saw_unpack_since_pack = True + self._remove_record(record) + previous_id = self.previous_by_generation.get(record.generation) + if previous_id is not None: + self._record_transition(previous_id, record.predictor_id) + self.previous_by_generation[record.generation] = record.predictor_id + successor_id = self.successors.get(record.predictor_id) + if successor_id and self.prefetch_allowed(): + self.prefetch_attempts += 1 + if not self.prefetch(record.generation, successor_id): + self.prefetch_misses += 1 + + def prefetch(self, generation: int, predictor_id: str) -> bool: + candidates = self.records.get((generation, predictor_id), ()) + for candidate in reversed(candidates): + if candidate.consumed or candidate.prefetched_tensor is not None: + continue + if candidate.original_device.type != "cuda" or not torch.cuda.is_available(): + self.prefetch_skipped += 1 + return False + if candidate.pool_key is None or not candidate.tensor.is_pinned(): + self.prefetch_skipped += 1 + return False + try: + copy_stream = _h2d_prefetch_stream_for(candidate.original_device) + with torch.cuda.stream(copy_stream): + if candidate.ready_event is not None: + copy_stream.wait_event(candidate.ready_event) + restored = candidate.tensor.to(candidate.original_device, non_blocking=True) + event = torch.cuda.Event() + event.record(copy_stream) + candidate.prefetched_tensor = restored + candidate.prefetch_event = event + self.prefetch_enqueued += 1 + return True + except RuntimeError: + candidate.prefetched_tensor = None + candidate.prefetch_event = None + self.prefetch_stale += 1 + return False + self.prefetch_skipped += 1 + return False + + def _remove_record(self, record: _OffloadedActivationRecord) -> None: + key = (record.generation, record.predictor_id) + entries = self.records.get(key) + if not entries: + return + try: + entries.remove(record) + except ValueError: + return + if not entries: + self.records.pop(key, None) + + def _retire_generation(self, generation: int) -> None: + stale_keys = [key for key in self.records if key[0] == generation] + for key in stale_keys: + entries = self.records.pop(key, ()) + for record in entries: + self._release_unused_record(record) + + def _release_unused_record(self, record: _OffloadedActivationRecord) -> None: + if record.cpu_released or record.pool_key is None: + return + release_event = record.prefetch_event or record.ready_event + _PINNED_MEMORY_POOL.release_after_record_event(record.pool_key, record.tensor, release_event) + record.cpu_released = True + record.prefetched_tensor = None + record.prefetch_event = None + + def _record_transition(self, previous_id: str, actual_id: str) -> None: + if self.disabled.get(previous_id): + return + counts = self.transitions[previous_id] + counts[actual_id] = int(counts.get(actual_id, 0)) + 1 + total = sum(int(value) for value in counts.values()) + best_id, best_count = max(counts.items(), key=lambda item: int(item[1])) + if total < _ACTIVATION_PREFETCH_MIN_OBSERVATIONS: + return + confidence = best_count / total if total else 0.0 + old_successor = self.successors.get(previous_id) + if confidence >= _ACTIVATION_PREFETCH_MIN_CONFIDENCE: + self.successors[previous_id] = best_id + if old_successor != best_id: + self.transition_updates += 1 + else: + self.successors.pop(previous_id, None) + + def record_hit(self) -> None: + self.prefetch_hits += 1 + + def prefetch_allowed(self) -> bool: + return _ACTIVATION_PREFETCH_ENABLED and not self.autotune_disabled + + @staticmethod + def _median(values: list[float]) -> float: + ordered = sorted(values) + midpoint = len(ordered) // 2 + if len(ordered) % 2: + return ordered[midpoint] + return (ordered[midpoint - 1] + ordered[midpoint]) / 2.0 + + def record_jit_restore_ms(self, value: float) -> None: + if not _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED or self.autotune_decision is not None: + return + self.jit_restore_ms.append(float(value)) + + def record_prefetch_wait_ms(self, value: float) -> None: + if not _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED or self.autotune_decision is not None: + return + self.prefetch_wait_ms.append(float(value)) + self._maybe_autotune() + + def _maybe_autotune(self) -> None: + if len(self.jit_restore_ms) < _ACTIVATION_PREFETCH_AUTOTUNE_MIN_SAMPLES: + return + if len(self.prefetch_wait_ms) < _ACTIVATION_PREFETCH_AUTOTUNE_MIN_SAMPLES: + return + jit_median = self._median(self.jit_restore_ms) + prefetch_median = self._median(self.prefetch_wait_ms) + if prefetch_median < jit_median * _ACTIVATION_PREFETCH_AUTOTUNE_MARGIN: + self.autotune_decision = "prefetch" + return + self.autotune_disabled = True + self.autotune_decision = "jit" + + def snapshot(self) -> dict: + active_records = sum(len(entries) for entries in self.records.values()) + return { + "enabled": _ACTIVATION_PREFETCH_ENABLED, + "active_records": active_records, + "generations_tracked": len({generation for generation, _logical_id in self.records}), + "total_packs": self.total_packs, + "total_unpacks": self.total_unpacks, + "prefetch_attempts": self.prefetch_attempts, + "prefetch_hits": self.prefetch_hits, + "prefetch_misses": self.prefetch_misses, + "prefetch_enqueued": self.prefetch_enqueued, + "prefetch_skipped": self.prefetch_skipped, + "prefetch_stale": self.prefetch_stale, + "transition_updates": self.transition_updates, + "learned_successors": len(self.successors), + "hit_rate": self.prefetch_hits / self.prefetch_attempts if self.prefetch_attempts else 0.0, + "autotune_enabled": _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED, + "autotune_disabled": self.autotune_disabled, + "autotune_decision": self.autotune_decision, + "jit_restore_samples": len(self.jit_restore_ms), + "prefetch_wait_samples": len(self.prefetch_wait_ms), + "jit_restore_median_ms": self._median(self.jit_restore_ms) if self.jit_restore_ms else None, + "prefetch_wait_median_ms": self._median(self.prefetch_wait_ms) if self.prefetch_wait_ms else None, + "copy_streams": get_activation_offload_copy_stream_stats(), + } + + def reset(self) -> None: + for entries in list(self.records.values()): + for record in entries: + self._release_unused_record(record) + self.records.clear() + self.transitions.clear() + self.successors.clear() + self.disabled.clear() + self.previous_by_generation.clear() + self.current_generation = 0 + self.saw_unpack_since_pack = False + self.total_packs = 0 + self.total_unpacks = 0 + self.prefetch_attempts = 0 + self.prefetch_hits = 0 + self.prefetch_misses = 0 + self.prefetch_enqueued = 0 + self.prefetch_skipped = 0 + self.prefetch_stale = 0 + self.transition_updates = 0 + self.jit_restore_ms.clear() + self.prefetch_wait_ms.clear() + self.autotune_disabled = False + self.autotune_decision = None + + +_ACTIVATION_PREFETCH_RUNTIME = _ActivationOffloadPrefetchRuntime() + + +def _device_index(device: torch.device) -> int: + return device.index if device.index is not None else torch.cuda.current_device() + + +def _d2h_copy_stream_for(device: torch.device) -> torch.cuda.Stream: + return _D2H_COPY_STREAMS.next(device) + + +def _h2d_prefetch_stream_for(device: torch.device) -> torch.cuda.Stream: + return _H2D_PREFETCH_STREAMS.next(device) + + +def set_activation_offload_d2h_copy_stream_count(count: int) -> None: + """Set the number of round-robin CUDA streams used for GPU-to-CPU offload copies.""" + _D2H_COPY_STREAMS.set_width(count) + + +def get_activation_offload_d2h_copy_stream_count() -> int: + return _D2H_COPY_STREAMS.width + + +def set_activation_offload_h2d_prefetch_stream_count(count: int) -> None: + """Set the number of round-robin CUDA streams used for CPU-to-GPU prefetch copies.""" + _H2D_PREFETCH_STREAMS.set_width(count) + + +def get_activation_offload_h2d_prefetch_stream_count() -> int: + return _H2D_PREFETCH_STREAMS.width + + +def get_activation_offload_copy_stream_stats() -> dict: + return { + "d2h": _D2H_COPY_STREAMS.snapshot(), + "h2d_prefetch": _H2D_PREFETCH_STREAMS.snapshot(), + } + + +def reset_activation_offload_copy_stream_stats() -> None: + _D2H_COPY_STREAMS.reset_stats() + _H2D_PREFETCH_STREAMS.reset_stats() + + +def set_activation_offload_pin_memory_max_buckets(max_buckets: int) -> None: + """Set the max number of distinct pinned CPU tensor buckets used for activation offload.""" + _PINNED_MEMORY_POOL.set_max_buckets(max_buckets) + + +def get_activation_offload_pin_memory_max_buckets() -> int: + return _PINNED_MEMORY_POOL.max_buckets + + +def normalize_activation_offload_pin_memory_max_buckets( + raw_value: object, default: int = _DEFAULT_PIN_MEMORY_MAX_BUCKETS +) -> int: + if raw_value in (None, "", "None"): + return default + try: + max_buckets = int(raw_value) + except (TypeError, ValueError) as exc: + raise ValueError("Gradient checkpointing offload pinned bucket count must be a non-negative integer.") from exc + if max_buckets < 0: + raise ValueError("Gradient checkpointing offload pinned bucket count must be non-negative.") + return max_buckets + + +def get_activation_offload_pin_memory_stats() -> dict: + """Return lifetime pinned CPU bucket statistics for activation offload.""" + return _PINNED_MEMORY_POOL.snapshot() + + +def reset_activation_offload_pin_memory_stats() -> None: + """Reset pinned CPU bucket statistics without clearing resident pinned buffers.""" + _PINNED_MEMORY_POOL.reset_stats() + + +def set_activation_offload_prefetch_enabled(enabled: bool) -> None: + """Enable learned H2D prefetching for labeled activation offload contexts.""" + global _ACTIVATION_PREFETCH_ENABLED + _ACTIVATION_PREFETCH_ENABLED = bool(enabled) + + +def get_activation_offload_prefetch_enabled() -> bool: + return _ACTIVATION_PREFETCH_ENABLED + + +def set_activation_offload_prefetch_autotune_enabled(enabled: bool) -> None: + """Enable first-run latency gating for activation prefetch.""" + global _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED + _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED = bool(enabled) + + +def get_activation_offload_prefetch_autotune_enabled() -> bool: + return _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED + + +def get_activation_offload_prefetch_stats() -> dict: + return _ACTIVATION_PREFETCH_RUNTIME.snapshot() + + +def reset_activation_offload_prefetch_stats() -> None: + _ACTIVATION_PREFETCH_RUNTIME.reset() + reset_activation_offload_copy_stream_stats() + + +def set_activation_offload_prefetch_runtime_disabled(disabled: bool, *, decision: str | None = None) -> None: + """Force the learned prefetch runtime on or off without changing payload labeling.""" + _ACTIVATION_PREFETCH_RUNTIME.autotune_disabled = bool(disabled) + _ACTIVATION_PREFETCH_RUNTIME.autotune_decision = decision + + +def mark_activation_offload_prefetch_autotune_decision(decision: str) -> None: + """Record the end-to-end autotune decision made by the trainer.""" + normalised = str(decision).lower() + if normalised not in {"prefetch", "jit"}: + raise ValueError("activation offload prefetch decision must be 'prefetch' or 'jit'") + set_activation_offload_prefetch_runtime_disabled(normalised == "jit", decision=normalised) + + +def activation_offload_prefetch_autotune_decision() -> str | None: + return _ACTIVATION_PREFETCH_RUNTIME.autotune_decision + class CPUOffloadHooks: """Context manager hooks that offload saved tensors to CPU during checkpointing.""" - def __init__(self): - # No global device state; device is tracked per saved tensor via pack/unpack payloads. - pass + def __init__( + self, + *, + offload_leaf_tensors: bool = False, + pin_memory: bool = True, + label: str | None = None, + prefetch: bool | None = None, + ): + self.offload_leaf_tensors = offload_leaf_tensors + self.pin_memory = pin_memory + self.label = str(label) if label else None + self.prefetch = _ACTIVATION_PREFETCH_ENABLED if prefetch is None else bool(prefetch) + self.pack_index = 0 + + @staticmethod + def _flat_storage_view(tensor: torch.Tensor) -> torch.Tensor | None: + if tensor.layout != torch.strided: + return None + span = 1 + sum((size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride(), strict=True)) + if span != tensor.numel(): + return None + return torch.as_strided(tensor, (tensor.numel(),), (1,), tensor.storage_offset()) + + def _transfer_view(self, tensor: torch.Tensor) -> tuple[torch.Tensor, _RestoreView | None]: + flat_view = self._flat_storage_view(tensor) + if flat_view is None: + return tensor, None + restore_storage_offset = tensor.storage_offset() - flat_view.storage_offset() + return flat_view, _RestoreView(tuple(tensor.size()), tuple(tensor.stride()), restore_storage_offset) + + def _copy_to_cpu( + self, tensor: torch.Tensor + ) -> tuple[torch.Tensor, _PinnedBucketKey | None, _RestoreView | None, torch.cuda.Event | None]: + detached = tensor.detach() + transfer_tensor, restore_view = self._transfer_view(detached) + if self.pin_memory: + key = _PINNED_MEMORY_POOL.key_for(transfer_tensor) + if key is not None: + cpu_tensor = _PINNED_MEMORY_POOL.checkout(key) + if cpu_tensor is not None: + try: + copy_stream = _d2h_copy_stream_for(tensor.device) + current_stream = torch.cuda.current_stream(tensor.device) + with torch.cuda.stream(copy_stream): + copy_stream.wait_stream(current_stream) + transfer_tensor.record_stream(copy_stream) + cpu_tensor.copy_(transfer_tensor, non_blocking=True) + ready_event = torch.cuda.Event() + ready_event.record(copy_stream) + return cpu_tensor, key, restore_view, ready_event + except RuntimeError: + _PINNED_MEMORY_POOL.discard_empty_bucket(key) + return transfer_tensor.to("cpu", non_blocking=True), None, restore_view, None def pack(self, tensor: torch.Tensor): """Called when a tensor is saved for backward - offload to CPU. @@ -25,33 +743,147 @@ def pack(self, tensor: torch.Tensor): restore only tensors that were actually offloaded, and to their correct original devices. """ - if tensor.device.type == "cuda": - cpu_tensor = tensor.to("cpu", non_blocking=True) - return cpu_tensor, tensor.device - # Tensor is already on CPU (or a non-CUDA device); mark as not offloaded. + if tensor.device.type == "cuda" and (self.offload_leaf_tensors or not tensor.is_leaf): + cpu_tensor, pool_key, restore_view, ready_event = self._copy_to_cpu(tensor) + if self.prefetch and self.label is not None: + generation = _ACTIVATION_PREFETCH_RUNTIME.next_generation_for_pack() + predictor_id = f"{self.label}:{self.pack_index}" + logical_id = f"{generation}:{predictor_id}:{id(cpu_tensor)}" + self.pack_index += 1 + record = _OffloadedActivationRecord( + logical_id=logical_id, + predictor_id=predictor_id, + generation=generation, + tensor=cpu_tensor, + original_device=tensor.device, + pool_key=pool_key, + restore_view=restore_view, + ready_event=ready_event, + ) + _ACTIVATION_PREFETCH_RUNTIME.register(record) + return record + return cpu_tensor, tensor.device, pool_key, restore_view, ready_event return tensor, None def unpack(self, payload) -> torch.Tensor: """Called when a tensor is needed for backward - restore to original device if needed.""" - # Expect payload of the form (cpu_tensor, original_device). + if isinstance(payload, _OffloadedActivationRecord): + return self._unpack_record(payload) + + # Expect payload of the form (cpu_tensor, original_device[, pool_key, restore_view]). try: - tensor, original_device = payload + tensor, original_device, *rest = payload except (TypeError, ValueError): # Fallback: if payload is not in the expected form, return it as-is. return payload if original_device is not None and tensor.device.type == "cpu": - return tensor.to(original_device, non_blocking=True) + current_stream = None + ready_event = rest[2] if len(rest) > 2 else None + if ready_event is not None and original_device.type == "cuda": + current_stream = torch.cuda.current_stream(original_device) + current_stream.wait_event(ready_event) + restored = tensor.to(original_device, non_blocking=True) + if original_device.type == "cuda": + current_stream = current_stream or torch.cuda.current_stream(original_device) + restored.record_stream(current_stream) + pool_key = rest[0] if rest else None + restore_view = rest[1] if len(rest) > 1 else None + if pool_key is not None: + _PINNED_MEMORY_POOL.release_after_cuda_copy(pool_key, tensor, original_device) + if restore_view is not None: + restored = torch.as_strided(restored, restore_view.size, restore_view.stride, restore_view.storage_offset) + return restored return tensor + def _release_cpu_tensor(self, record: _OffloadedActivationRecord) -> None: + if record.cpu_released or record.pool_key is None: + return + _PINNED_MEMORY_POOL.release_after_cuda_copy(record.pool_key, record.tensor, record.original_device) + record.cpu_released = True + + def _restore_view_if_needed(self, tensor: torch.Tensor, restore_view: _RestoreView | None) -> torch.Tensor: + if restore_view is None: + return tensor + return torch.as_strided(tensor, restore_view.size, restore_view.stride, restore_view.storage_offset) + + def _unpack_record(self, record: _OffloadedActivationRecord) -> torch.Tensor: + record.consumed = True + if record.prefetched_tensor is not None: + elapsed_ms = None + if record.prefetch_event is not None and record.original_device.type == "cuda": + current_stream = torch.cuda.current_stream(record.original_device) + if _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED and _ACTIVATION_PREFETCH_RUNTIME.autotune_decision is None: + start_event = torch.cuda.Event(enable_timing=True) + end_event = torch.cuda.Event(enable_timing=True) + start_event.record(current_stream) + current_stream.wait_event(record.prefetch_event) + end_event.record(current_stream) + end_event.synchronize() + elapsed_ms = start_event.elapsed_time(end_event) + else: + current_stream.wait_event(record.prefetch_event) + record.prefetched_tensor.record_stream(current_stream) + restored = record.prefetched_tensor + _ACTIVATION_PREFETCH_RUNTIME.record_hit() + if elapsed_ms is not None: + _ACTIVATION_PREFETCH_RUNTIME.record_prefetch_wait_ms(elapsed_ms) + self._release_cpu_tensor(record) + _ACTIVATION_PREFETCH_RUNTIME.consume(record) + return self._restore_view_if_needed(restored, record.restore_view) + + if record.ready_event is not None and record.original_device.type == "cuda": + torch.cuda.current_stream(record.original_device).wait_event(record.ready_event) + elapsed_ms = None + if ( + _ACTIVATION_PREFETCH_AUTOTUNE_ENABLED + and _ACTIVATION_PREFETCH_RUNTIME.autotune_decision is None + and record.original_device.type == "cuda" + ): + current_stream = torch.cuda.current_stream(record.original_device) + start_event = torch.cuda.Event(enable_timing=True) + end_event = torch.cuda.Event(enable_timing=True) + start_event.record(current_stream) + restored = record.tensor.to(record.original_device, non_blocking=True) + end_event.record(current_stream) + end_event.synchronize() + elapsed_ms = start_event.elapsed_time(end_event) + else: + restored = record.tensor.to(record.original_device, non_blocking=True) + if record.original_device.type == "cuda": + restored.record_stream(torch.cuda.current_stream(record.original_device)) + if elapsed_ms is not None: + _ACTIVATION_PREFETCH_RUNTIME.record_jit_restore_ms(elapsed_ms) + self._release_cpu_tensor(record) + _ACTIVATION_PREFETCH_RUNTIME.consume(record) + return self._restore_view_if_needed(restored, record.restore_view) + + +def activation_offload_context(enabled: bool = True, *, label: str | None = None, prefetch: bool | None = None): + """Offload non-leaf CUDA tensors saved for backward inside the context.""" + if not enabled: + return nullcontext() + if label is not None: + frame = sys._getframe(1) + label = f"{label}:{frame.f_code.co_name}:{frame.f_lineno}" + hooks = CPUOffloadHooks(offload_leaf_tensors=False, label=label, prefetch=prefetch) + return torch.autograd.graph.saved_tensors_hooks(hooks.pack, hooks.unpack) + + +def activation_offload(function, *args, **kwargs): + """Run a function normally while offloading saved non-leaf CUDA tensors.""" + kwargs.pop("use_reentrant", None) + with activation_offload_context(): + return function(*args, **kwargs) + def offloaded_checkpoint(function, *args, use_reentrant: bool = False, **kwargs): """ Drop-in replacement for torch.utils.checkpoint.checkpoint using CPU offload. - Instead of recomputing activations during backward pass (standard checkpointing), - this offloads saved tensors to CPU asynchronously and restores them when needed. - This approach trades PCIe bandwidth for GPU memory savings. + This still uses PyTorch checkpoint rematerialization. Saved tensor hooks + offload tensors that autograd saves inside the checkpointed region and + restore them when needed. Args: function: The forward function to checkpoint diff --git a/simpletuner/helpers/training/save_hooks.py b/simpletuner/helpers/training/save_hooks.py index 2d334e2e2..387b28435 100644 --- a/simpletuner/helpers/training/save_hooks.py +++ b/simpletuner/helpers/training/save_hooks.py @@ -300,9 +300,30 @@ def __init__( self.denoiser_class = self.model.MODEL_CLASS self.denoiser_subdir = self.model.MODEL_SUBFOLDER - self.pipeline_class = self.model.PIPELINE_CLASSES[ - (PipelineTypes.IMG2IMG if args.validation_using_datasets else PipelineTypes.TEXT2IMG) - ] + pipeline_type = getattr(self.model, "DEFAULT_PIPELINE_TYPE", PipelineTypes.TEXT2IMG) + if isinstance(pipeline_type, str): + try: + pipeline_type = PipelineTypes(pipeline_type) + except ValueError as exc: + raise ValueError( + f"Unsupported DEFAULT_PIPELINE_TYPE for {type(self.model).__name__}: {pipeline_type!r}" + ) from exc + if args.validation_using_datasets and PipelineTypes.IMG2IMG in self.model.PIPELINE_CLASSES: + pipeline_type = PipelineTypes.IMG2IMG + if not isinstance(pipeline_type, PipelineTypes): + raise ValueError( + f"DEFAULT_PIPELINE_TYPE for {type(self.model).__name__} must be a PipelineTypes value, " + f"got {pipeline_type!r}." + ) + self.pipeline_class = self.model.PIPELINE_CLASSES.get(pipeline_type) + if self.pipeline_class is None: + available_pipeline_types = ", ".join( + key.value if isinstance(key, PipelineTypes) else repr(key) for key in self.model.PIPELINE_CLASSES + ) + raise ValueError( + f"{type(self.model).__name__} does not register a save pipeline for {pipeline_type.value!r}. " + f"Available pipeline types: {available_pipeline_types or 'none'}." + ) self.ema_model_cls = self.model.get_trained_component().__class__ self.ema_model_subdir = f"{self.model.MODEL_SUBFOLDER}_ema" @@ -721,10 +742,7 @@ def _save_lora(self, models, weights, output_dir, write: bool = True): self.ema_model.copy_to(trainable_parameters) ema_trained_component = unwrap_model(self.accelerator, self.model.get_trained_component()) lora_save_parameters = { - f"{self.model.MODEL_SUBFOLDER}_lora_layers": convert_state_dict_to_diffusers( - get_peft_model_state_dict(ema_trained_component), - original_type=StateDictType.PEFT, - ), + f"{self.model.MODEL_SUBFOLDER}_lora_layers": get_peft_model_state_dict(ema_trained_component), } ema_modules_to_save = {self.model.MODEL_SUBFOLDER: ema_trained_component} ema_metadata = _collate_lora_metadata(ema_modules_to_save) @@ -771,9 +789,8 @@ def _save_lora(self, models, weights, output_dir, write: bool = True): modules_to_save["controlnet"] = unwrapped_model elif isinstance(unwrapped_model, tuple(trained_component_classes)): # unet_lora_layers or transformer_lora_layers - lora_save_parameters[f"{self.model.MODEL_SUBFOLDER}_lora_layers"] = convert_state_dict_to_diffusers( - get_peft_model_state_dict(unwrapped_model), - original_type=StateDictType.PEFT, + lora_save_parameters[f"{self.model.MODEL_SUBFOLDER}_lora_layers"] = get_peft_model_state_dict( + unwrapped_model ) modules_to_save[self.model.MODEL_SUBFOLDER] = unwrapped_model elif text_encoder_0_cls is not None and isinstance(unwrapped_model, text_encoder_0_cls): @@ -1071,13 +1088,20 @@ def _save_full_model( def save_model_hook(self, models, weights, output_dir): # Write "training_state.json" to the output directory containing the training state StateTracker.save_training_state(os.path.join(output_dir, self.training_state_path)) - StateTracker.save_ramtorch_prefetch_orders(output_dir) distributed_type = DistributedType.NO is_main_process = True + is_local_main_process = True if self.accelerator is not None: distributed_type = getattr(self.accelerator, "distributed_type", DistributedType.NO) is_main_process = getattr(self.accelerator, "is_main_process", True) + is_local_main_process = getattr(self.accelerator, "is_local_main_process", True) + + # One writer per node, so the file also lands on node-local storage in + # multi-node runs. Concurrent node-mains on a shared filesystem are safe + # because the writer uses a process-unique temp name and atomic rename. + if is_local_main_process: + StateTracker.save_ramtorch_prefetch_orders(output_dir) with self._offload_models_during_save(is_main_process): self._save_ema_state(output_dir, is_main_process=is_main_process) diff --git a/simpletuner/helpers/training/state_tracker.py b/simpletuner/helpers/training/state_tracker.py index 076ce752e..eeb4b6f3f 100644 --- a/simpletuner/helpers/training/state_tracker.py +++ b/simpletuner/helpers/training/state_tracker.py @@ -405,7 +405,7 @@ def save_ramtorch_prefetch_orders(cls, directory: str | Path | None = None) -> N return data = cls._normalise_ramtorch_prefetch_orders(cls.ramtorch_prefetch_orders) path.parent.mkdir(parents=True, exist_ok=True) - temp_path = path.with_suffix(f"{path.suffix}.tmp") + temp_path = path.with_suffix(f"{path.suffix}.tmp.{os.getpid()}") with temp_path.open("w") as handle: fcntl.flock(handle, fcntl.LOCK_EX) try: diff --git a/simpletuner/helpers/training/trainer.py b/simpletuner/helpers/training/trainer.py index 36ea7a6fe..42d01b29c 100644 --- a/simpletuner/helpers/training/trainer.py +++ b/simpletuner/helpers/training/trainer.py @@ -3386,6 +3386,9 @@ def enable_gradient_checkpointing(self): logger.debug("Enabling gradient checkpointing.") if hasattr(self.model.get_trained_component(), "enable_gradient_checkpointing"): unwrap_model(self.accelerator, self.model.get_trained_component()).enable_gradient_checkpointing() + model_level_enable = getattr(self.model, "enable_gradient_checkpointing", None) + if callable(model_level_enable): + model_level_enable() if hasattr(self.config, "train_text_encoder") and self.config.train_text_encoder: for text_encoder in self.model.text_encoders: if text_encoder is not None: @@ -3398,6 +3401,9 @@ def disable_gradient_checkpointing(self): unwrap_model( self.accelerator, self.model.get_trained_component(base_model=True) ).disable_gradient_checkpointing() + model_level_disable = getattr(self.model, "disable_gradient_checkpointing", None) + if callable(model_level_disable): + model_level_disable() if self.config.controlnet: unwrap_model(self.accelerator, self.model.get_trained_component()).disable_gradient_checkpointing() if hasattr(self.config, "train_text_encoder") and self.config.train_text_encoder: @@ -4172,6 +4178,9 @@ def init_prepare_models(self, lr_scheduler): attach_shared_ramtorch_parameters as attach_shared_ramtorch_parameters, ) + moved = ramtorch_utils.move_embeddings_to_device(primary_model, self.accelerator.device) + if moved: + logger.info("Moved %s non-RamTorch CPU parameters/buffers before DDP prepare.", moved) ignored = ramtorch_utils.mark_ddp_ignore_params(primary_model) if ignored: logger.info("Marking %s RamTorch parameters to ignore for DDP.", ignored) @@ -5241,15 +5250,21 @@ def _run_standard_checkpoint(self, webhook_message: str | None, parent_loss, epo **progress_kwargs, ) self._emit_event(event) - if self.accelerator.is_main_process and self.config.checkpoints_total_limit is not None: - self.checkpoint_state_cleanup( - self.config.output_dir, - self.config.checkpoints_total_limit, - ) + checkpoint_limit = self.config.checkpoints_total_limit + if self.accelerator.is_main_process: + self.checkpoint_state_cleanup_temp(self.config.output_dir) save_path = None if self.accelerator.is_main_process or self.config.use_deepspeed_optimizer or self.config.fsdp_enable: save_path = self.checkpoint_state_save(self.config.output_dir) + if self.accelerator.is_main_process and checkpoint_limit is not None and checkpoint_limit > 0: + if len(self.checkpoint_state_filter(self.config.output_dir)) > checkpoint_limit: + self._drain_hub_upload_futures(wait=True) + self.checkpoint_state_cleanup( + self.config.output_dir, + checkpoint_limit, + protected_checkpoint=save_path, + ) hub_upload_planned = upload_to_hub and self.hub_manager is not None if hub_upload_planned: @@ -5264,6 +5279,7 @@ def _upload_latest_checkpoint(): webhook_handler=self.webhook_handler, global_step=captured_step, epoch=captured_epoch, + checkpoint_path=save_path, ) return remote_path, local_path, repo_url @@ -5277,6 +5293,23 @@ def _upload_latest_checkpoint(): self._run_post_upload_script(local_path=save_path, remote_path=None) return save_path + def _save_rolling_checkpoint(self): + checkpoint_limit = self.config.checkpoints_rolling_total_limit + if self.accelerator.is_main_process: + self.checkpoint_state_cleanup_temp(self.config.output_dir) + + save_path = None + if self.accelerator.is_main_process or self.config.use_deepspeed_optimizer or self.config.fsdp_enable: + save_path = self.checkpoint_state_save(self.config.output_dir, "rolling") + if self.accelerator.is_main_process and checkpoint_limit is not None and checkpoint_limit > 0: + self.checkpoint_state_cleanup( + self.config.output_dir, + checkpoint_limit, + "rolling", + protected_checkpoint=save_path, + ) + return save_path + def _send_webhook_msg( self, message: str, message_level: str = "info", store_response: bool = False, emit_event: bool = True ): @@ -5501,7 +5534,8 @@ def _epoch_rollover(self, epoch): # we'll have to split the buckets between GPUs again now, so that the VAE cache distributes properly. logger.info("Splitting buckets across GPUs") backend["metadata_backend"].split_buckets_between_processes( - gradient_accumulation_steps=self.config.gradient_accumulation_steps + gradient_accumulation_steps=self.config.gradient_accumulation_steps, + apply_padding=(self.config.overrode_max_train_steps or self.config.allow_dataset_oversubscription), ) # we have to rebuild the VAE cache if it exists. if "vaecache" in backend: @@ -5640,6 +5674,119 @@ def model_predict( return model_pred + def _compute_model_prediction_loss(self, prepared_batch: dict) -> tuple[torch.Tensor, dict, torch.Tensor, dict, dict]: + model_pred = self.model_predict( + prepared_batch=prepared_batch, + ) + loss, loss_logs = self.model.loss_with_logs( + prepared_batch=prepared_batch, + model_output=model_pred, + apply_conditioning_mask=True, + ) + diffusion_loss = loss.clone() + loss, aux_loss_logs = self.model.auxiliary_loss( + prepared_batch=prepared_batch, + model_output=model_pred, + loss=loss, + ) + distill_logs = {} + if self.config.distillation_method is not None: + loss, distill_logs = self.distiller.compute_distill_loss(prepared_batch, model_pred, loss) + loss, gen_logs = self.distiller.generator_loss_step(prepared_batch, model_pred, loss) + distill_logs.update(gen_logs) + return loss, loss_logs, diffusion_loss, aux_loss_logs, distill_logs + + def _discard_probe_gradients(self) -> None: + trained_component = self.model.get_trained_component(unwrap_model=False) + if trained_component is not None: + trained_component.zero_grad(set_to_none=True) + if getattr(self, "optimizer", None) is not None: + self.optimizer.zero_grad(set_to_none=True) + if getattr(self, "sidecar_optimizer", None) is not None: + self.sidecar_optimizer.zero_grad(set_to_none=True) + + def _activation_offload_prefetch_autotune_probe(self, prepared_batch: dict, *, label: str) -> float: + if torch.cuda.is_available(): + torch.cuda.synchronize() + start_time = time.perf_counter() + loss, _loss_logs, _diffusion_loss, _aux_loss_logs, _distill_logs = self._compute_model_prediction_loss( + dict(prepared_batch) + ) + self.accelerator.backward(loss) + if torch.cuda.is_available(): + torch.cuda.synchronize() + elapsed = time.perf_counter() - start_time + self._discard_probe_gradients() + logger.debug("Activation offload prefetch autotune %s probe: %.4fs", label, elapsed) + return elapsed + + def _maybe_autotune_activation_offload_prefetch(self, prepared_batch: dict) -> None: + if not getattr(self.config, "gradient_checkpointing_offload_prefetch", False): + return + if not getattr(self.config, "gradient_checkpointing_offload_attention", False): + logger.info("Skipping activation offload prefetch autotune because attention activation offload is disabled.") + return + if getattr(self, "_activation_offload_prefetch_autotuned", False): + return + if self.config.disable_accelerator or not torch.cuda.is_available(): + logger.info("Skipping activation offload prefetch autotune because CUDA training is unavailable.") + return + if self.config.distillation_method is not None: + logger.warning("Skipping activation offload prefetch autotune for distillation runs.") + self._activation_offload_prefetch_autotuned = True + return + + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + get_activation_offload_prefetch_autotune_enabled, + mark_activation_offload_prefetch_autotune_decision, + reset_activation_offload_prefetch_stats, + set_activation_offload_prefetch_autotune_enabled, + set_activation_offload_prefetch_enabled, + set_activation_offload_prefetch_runtime_disabled, + ) + + self._activation_offload_prefetch_autotuned = True + previous_autotune_enabled = get_activation_offload_prefetch_autotune_enabled() + cpu_rng_state = torch.get_rng_state() + cuda_rng_states = torch.cuda.get_rng_state_all() + + def restore_rng_state() -> None: + torch.set_rng_state(cpu_rng_state) + torch.cuda.set_rng_state_all(cuda_rng_states) + + logger.info("Autotuning attention activation prefetch with a real model_predict/backward probe.") + try: + set_activation_offload_prefetch_enabled(True) + set_activation_offload_prefetch_autotune_enabled(False) + reset_activation_offload_prefetch_stats() + + set_activation_offload_prefetch_runtime_disabled(True, decision=None) + restore_rng_state() + self._activation_offload_prefetch_autotune_probe(prepared_batch, label="learn") + + restore_rng_state() + jit_seconds = self._activation_offload_prefetch_autotune_probe(prepared_batch, label="jit") + + set_activation_offload_prefetch_runtime_disabled(False, decision=None) + restore_rng_state() + prefetch_seconds = self._activation_offload_prefetch_autotune_probe(prepared_batch, label="prefetch") + + decision = "prefetch" if prefetch_seconds < jit_seconds * 0.98 else "jit" + mark_activation_offload_prefetch_autotune_decision(decision) + logger.info( + "Activation offload prefetch autotune selected %s (jit=%.4fs, prefetch=%.4fs).", + decision, + jit_seconds, + prefetch_seconds, + ) + except Exception as exc: + logger.warning("Activation offload prefetch autotune failed; disabling prefetch. Error: %s", exc) + mark_activation_offload_prefetch_autotune_decision("jit") + self._discard_probe_gradients() + finally: + restore_rng_state() + set_activation_offload_prefetch_autotune_enabled(previous_autotune_enabled) + def _prepare_custom_timestep_batch(self, prepared_batch: dict, custom_timesteps): overridden_batch = dict(prepared_batch) reference_timesteps = prepared_batch.get("timesteps") @@ -5861,7 +6008,7 @@ def checkpoint_state_filter(self, output_dir, suffix=None): if len(cs) < 2: continue elif len(cs) > 2: - sfx = cs[2] + sfx = cs[-1] if base != "checkpoint": continue @@ -5874,26 +6021,36 @@ def checkpoint_state_filter(self, output_dir, suffix=None): return checkpoints_keep - def checkpoint_state_cleanup(self, output_dir, limit, suffix=None): + def checkpoint_state_cleanup_temp(self, output_dir): if self.checkpoint_manager: - self.checkpoint_manager.cleanup_checkpoints(limit, suffix) + self.checkpoint_manager.cleanup_temp_checkpoints() else: - # Fallback to original implementation - # remove any left over temp checkpoints (partially written, etc) checkpoints = self.checkpoint_state_filter(output_dir, "tmp") for removing_checkpoint in checkpoints: self.checkpoint_state_remove(output_dir, removing_checkpoint) + def checkpoint_state_cleanup(self, output_dir, limit, suffix=None, protected_checkpoint=None): + if self.checkpoint_manager: + self.checkpoint_manager.cleanup_checkpoints( + limit, + suffix, + protected_checkpoint=protected_checkpoint, + ) + else: + # remove any left over temp checkpoints (partially written, etc) + self.checkpoint_state_cleanup_temp(output_dir) + # now remove normal checkpoints past the limit checkpoints = self.checkpoint_state_filter(output_dir, suffix) checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) - # before we save the new checkpoint, we need to have at _most_ `limit - 1` checkpoints - if len(checkpoints) < limit: + if len(checkpoints) <= limit: return - num_to_remove = len(checkpoints) - limit + 1 - removing_checkpoints = checkpoints[0:num_to_remove] + num_to_remove = len(checkpoints) - limit + protected_name = os.path.basename(os.path.normpath(protected_checkpoint)) if protected_checkpoint else None + removal_candidates = [checkpoint for checkpoint in checkpoints if checkpoint != protected_name] + removing_checkpoints = removal_candidates[0:num_to_remove] logger.debug(f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints") logger.debug(f"removing checkpoints: {', '.join(removing_checkpoints)}") @@ -5971,6 +6128,13 @@ def checkpoint_state_save(self, output_dir, suffix=None): all_processes_saving = bool( getattr(self.config, "use_deepspeed_optimizer", False) or getattr(self.config, "fsdp_enable", False) ) + if ( + self.accelerator is not None + and distributed_type != DistributedType.NO + and all_processes_saving + and self.config.checkpointing_use_tempdir + ): + self.accelerator.wait_for_everyone() if fsdp_v2_run: logger.info("FSDP v2 detected; saving with sharded state dict (_use_dtensor disabled for NCCL compatibility).") if is_main_process: @@ -6422,26 +6586,12 @@ def train(self): module.scaling[key] = val * strength slider_original_scaling[layer_id] = (module, saved) + self._maybe_autotune_activation_offload_prefetch(prepared_batch) + training_logger.debug("Predicting.") - model_pred = self.model_predict( - prepared_batch=prepared_batch, - ) - loss, loss_logs = self.model.loss_with_logs( - prepared_batch=prepared_batch, - model_output=model_pred, - apply_conditioning_mask=True, + loss, loss_logs, diffusion_loss, aux_loss_logs, distill_logs = self._compute_model_prediction_loss( + prepared_batch ) - diffusion_loss = loss.clone() - loss, aux_loss_logs = self.model.auxiliary_loss( - prepared_batch=prepared_batch, - model_output=model_pred, - loss=loss, - ) - distill_logs = {} - if self.config.distillation_method is not None: - loss, distill_logs = self.distiller.compute_distill_loss(prepared_batch, model_pred, loss) - loss, gen_logs = self.distiller.generator_loss_step(prepared_batch, model_pred, loss) - distill_logs.update(gen_logs) parent_loss = None if is_regularisation_data: parent_loss = loss @@ -6812,16 +6962,7 @@ def train(self): **progress_kwargs, ) self._emit_event(event) - if self.accelerator.is_main_process and self.config.checkpoints_rolling_total_limit is not None: - # _before_ saving state, check if this save would set us over the `checkpoints_rolling_total_limit` - self.checkpoint_state_cleanup( - self.config.output_dir, - self.config.checkpoints_rolling_total_limit, - "rolling", - ) - - if self.accelerator.is_main_process or self.config.use_deepspeed_optimizer: - self.checkpoint_state_save(self.config.output_dir, "rolling") + self._save_rolling_checkpoint() if ( self.config.accelerator_cache_clear_interval is not None diff --git a/simpletuner/helpers/training/validation.py b/simpletuner/helpers/training/validation.py index d2b52a5be..64ae9d0af 100644 --- a/simpletuner/helpers/training/validation.py +++ b/simpletuner/helpers/training/validation.py @@ -696,6 +696,26 @@ def retrieve_validation_s2v_samples() -> list[ValidationPrompt]: args = StateTracker.get_args() validation_set = [] + def resolve_video_sample_path(sample_path: str, backend_config: dict, sampler: Any) -> str: + if not sample_path: + return sample_path + if Path(sample_path).is_absolute(): + return sample_path + + sample_root = backend_config.get("instance_data_dir") + if not sample_root: + metadata_backend = getattr(sampler, "metadata_backend", None) + sample_root = getattr(metadata_backend, "instance_data_dir", None) + if not sample_root: + return sample_path + + return os.path.join(sample_root, sample_path.lstrip(os.sep)) + + def local_audio_path_or_none(audio_path: str | None) -> str | None: + if not audio_path: + return None + return audio_path if Path(audio_path).exists() else None + # Get video backends that have s2v_datasets linked video_backends = StateTracker.get_data_backends(_type="video") selected_eval_backend_ids = _assert_eval_dataset_exists(args.eval_dataset_id, video_backends, "video validation") @@ -732,15 +752,14 @@ def retrieve_validation_s2v_samples() -> list[ValidationPrompt]: # Find matching audio from s2v_datasets audio_path = None if sample_path is not None: - from pathlib import Path - + resolved_sample_path = resolve_video_sample_path(sample_path, backend_config, sampler) video_stem = Path(sample_path).stem for s2v_dataset in s2v_datasets: s2v_config = s2v_dataset.get("config", {}) audio_config = s2v_config.get("audio", {}) if audio_config.get("source_from_video", False): - audio_path = sample_path + audio_path = local_audio_path_or_none(resolved_sample_path) break audio_root = s2v_config.get("instance_data_dir") if not audio_root: @@ -791,12 +810,15 @@ def _validation_text_cache_key(args, shortname: str, prompt: str) -> str: def prepare_validation_prompt_list(args, embed_cache, model): - precompute_text_embeddings = not getattr(embed_cache, "text_cache_ondemand", False) + model_uses_text_cache = True + if hasattr(model, "uses_text_embeddings_cache"): + model_uses_text_cache = model.uses_text_embeddings_cache() + precompute_text_embeddings = model_uses_text_cache and not getattr(embed_cache, "text_cache_ondemand", False) validation_prompts: list[PromptLibraryEntry] = ( [PromptLibraryEntry(prompt="")] if not StateTracker.get_args().validation_disable_unconditional else [] ) validation_shortnames = ["unconditional"] if not StateTracker.get_args().validation_disable_unconditional else [] - if not hasattr(embed_cache, "model_type"): + if model_uses_text_cache and not hasattr(embed_cache, "model_type"): raise ValueError( f"The default text embed cache backend was not found. You must specify 'default: true' on your text embed data backend via {StateTracker.get_args().data_backend_config}." ) @@ -810,7 +832,7 @@ def prepare_validation_prompt_list(args, embed_cache, model): } if precompute_text_embeddings: embed_cache.compute_embeddings_for_prompts([prompt_record], is_validation=True, load_from_cache=False) - model_type = embed_cache.model_type + model_type = getattr(embed_cache, "model_type", None) validation_sample_images = None deepfloyd_stage2_needs_validation_images = ( "deepfloyd" in args.model_family @@ -2095,7 +2117,10 @@ def _discover_validation_input_samples(self): def _pipeline_cls(self): if self.model is not None: - if self.config.validation_using_datasets: + if isinstance(self.model, AudioModelFoundation): + if PipelineTypes.TEXT2AUDIO not in self.model.PIPELINE_CLASSES: + raise ValueError(f"Cannot run {self.model.MODEL_CLASS} in Text2Audio mode for validation.") + if self.config.validation_using_datasets and not isinstance(self.model, AudioModelFoundation): if PipelineTypes.IMG2IMG not in self.model.PIPELINE_CLASSES: raise ValueError(f"Cannot run {self.model.MODEL_CLASS} in Img2Img mode for validation.") if self.config.controlnet: @@ -2111,6 +2136,8 @@ def _pipeline_cls(self): return self.model.PIPELINE_CLASSES[PipelineTypes.CONTROLNET] if self.config.control: return self.model.PIPELINE_CLASSES[PipelineTypes.CONTROL] + if isinstance(self.model, AudioModelFoundation): + return self.model.PIPELINE_CLASSES[PipelineTypes.TEXT2AUDIO] return self.model.PIPELINE_CLASSES[PipelineTypes.TEXT2IMG] def _gather_prompt_embeds( @@ -2857,6 +2884,9 @@ def setup_pipeline(self, validation_type): if getattr(self.model, "requires_s2v_validation_inputs", lambda: False)(): if PipelineTypes.IMG2VIDEO in self.model.PIPELINE_CLASSES: pipeline_type = PipelineTypes.IMG2VIDEO + elif isinstance(self.model, AudioModelFoundation): + if PipelineTypes.TEXT2AUDIO in self.model.PIPELINE_CLASSES: + pipeline_type = PipelineTypes.TEXT2AUDIO elif self.config.validation_using_datasets: if PipelineTypes.IMG2IMG in self.model.PIPELINE_CLASSES: pipeline_type = PipelineTypes.IMG2IMG @@ -4333,6 +4363,26 @@ def validate_prompt( "image": validation_input_image_for_resolution, "audio_path": s2v_audio_path, } + if resolution[0] > 0 and resolution[1] > 0: + if isinstance(extra_validation_kwargs["image"], list): + extra_validation_kwargs["image"] = [ + img.resize(resolution, Image.Resampling.LANCZOS) + for img in extra_validation_kwargs["image"] + if isinstance(img, Image.Image) + ] + elif isinstance(extra_validation_kwargs["image"], Image.Image): + extra_validation_kwargs["image"] = extra_validation_kwargs["image"].resize( + resolution, Image.Resampling.LANCZOS + ) + validation_input_image_for_resolution = extra_validation_kwargs["image"] + extra_validation_kwargs["_s2v_conditioning"]["image"] = validation_input_image_for_resolution + validation_resolution_width, validation_resolution_height = resolution + elif isinstance(validation_input_image_for_resolution, list) and validation_input_image_for_resolution: + validation_resolution_width, validation_resolution_height = validation_input_image_for_resolution[0].size + elif isinstance(validation_input_image_for_resolution, Image.Image): + validation_resolution_width, validation_resolution_height = validation_input_image_for_resolution.size + else: + validation_resolution_width, validation_resolution_height = 0, 0 elif validation_input_image is not None: validation_input_image_for_resolution = _coerce_validation_image_input(validation_input_image) extra_validation_kwargs["image"] = validation_input_image_for_resolution diff --git a/simpletuner/helpers/utils/checkpoint_manager.py b/simpletuner/helpers/utils/checkpoint_manager.py index 79a9e7120..51e779b6b 100644 --- a/simpletuner/helpers/utils/checkpoint_manager.py +++ b/simpletuner/helpers/utils/checkpoint_manager.py @@ -187,12 +187,18 @@ def load_manifest(self, checkpoint_path: str) -> Optional[Dict[str, object]]: return None return data - def cleanup_checkpoints(self, limit: int, suffix: Optional[str] = None): - """Clean up old checkpoints, keeping only the most recent ones. + def cleanup_checkpoints( + self, + limit: int, + suffix: Optional[str] = None, + protected_checkpoint: Optional[str] = None, + ): + """Clean up old checkpoints, preserving a specified checkpoint when provided. Args: limit: Maximum number of checkpoints to keep suffix: Optional suffix to filter checkpoints + protected_checkpoint: Optional checkpoint path or name to preserve """ # Remove temp checkpoints first self._remove_temp_checkpoints() @@ -204,7 +210,9 @@ def cleanup_checkpoints(self, limit: int, suffix: Optional[str] = None): # Remove old checkpoints if we exceed the limit if len(checkpoints) > limit: num_to_remove = len(checkpoints) - limit - removing_checkpoints = checkpoints[:num_to_remove] + protected_name = os.path.basename(os.path.normpath(protected_checkpoint)) if protected_checkpoint else None + removal_candidates = [checkpoint for checkpoint in checkpoints if checkpoint != protected_name] + removing_checkpoints = removal_candidates[:num_to_remove] logger.debug(f"{len(checkpoints)} checkpoints exist, removing {len(removing_checkpoints)} checkpoints") logger.debug(f"Removing checkpoints: {', '.join(removing_checkpoints)}") @@ -212,6 +220,10 @@ def cleanup_checkpoints(self, limit: int, suffix: Optional[str] = None): for checkpoint in removing_checkpoints: self.remove_checkpoint(checkpoint) + def cleanup_temp_checkpoints(self): + """Remove temporary checkpoints without rotating completed checkpoints.""" + self._remove_temp_checkpoints() + def remove_checkpoint(self, checkpoint_name: str): """Remove a specific checkpoint. @@ -246,7 +258,7 @@ def _filter_checkpoints(self, suffix: Optional[str] = None) -> List[str]: if len(parts) < 2: continue elif len(parts) > 2: - checkpoint_suffix = parts[2] + checkpoint_suffix = parts[-1] if base != "checkpoint": continue diff --git a/simpletuner/helpers/utils/ramtorch.py b/simpletuner/helpers/utils/ramtorch.py index 7d81f4540..4539467c1 100644 --- a/simpletuner/helpers/utils/ramtorch.py +++ b/simpletuner/helpers/utils/ramtorch.py @@ -778,20 +778,26 @@ def move_non_linear_layers_to_device(module: nn.Module, device: object) -> int: def mark_ddp_ignore_params(module: nn.Module) -> int: """ - Mark RamTorch parameters on a module to be ignored by DistributedDataParallel. + Mark RamTorch/CPU-resident parameters on a module to be ignored by + DistributedDataParallel. Returns: - Number of parameters added to the ignore list. + Number of parameters and buffers added to the ignore list. """ - ramtorch_names = [name for name, param in module.named_parameters() if getattr(param, "is_ramtorch", False)] - ramtorch_names.extend(name for name, buffer in module.named_buffers() if getattr(buffer, "is_ramtorch", False)) - if not ramtorch_names: - device_types = {param.device.type for _, param in module.named_parameters() if param.device is not None} - if "cuda" in device_types and "cpu" in device_types: - ramtorch_names = [name for name, param in module.named_parameters() if param.device.type == "cpu"] - if not ramtorch_names: + named_params = list(module.named_parameters()) + named_buffers = list(module.named_buffers()) + ignore_names = {name for name, param in named_params if getattr(param, "is_ramtorch", False)} + ignore_names.update(name for name, buffer in named_buffers if getattr(buffer, "is_ramtorch", False)) + + device_types = {param.device.type for _, param in named_params if param.device is not None} + should_ignore_cpu_residents = bool(ignore_names) or {"cpu", "cuda"}.issubset(device_types) + if should_ignore_cpu_residents: + ignore_names.update(name for name, param in named_params if param.device.type == "cpu" and not param.requires_grad) + ignore_names.update(name for name, buffer in named_buffers if buffer.device.type == "cpu") + if not ignore_names: return 0 - existing = getattr(module, "_ddp_params_and_buffers_to_ignore", set()) - module._ddp_params_and_buffers_to_ignore = set(existing) | set(ramtorch_names) - return len(ramtorch_names) + existing = set(getattr(module, "_ddp_params_and_buffers_to_ignore", set())) + newly_added = ignore_names - existing + module._ddp_params_and_buffers_to_ignore = existing | ignore_names + return len(newly_added) diff --git a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py index dafe9f727..70283c63a 100644 --- a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py +++ b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/model.py @@ -490,14 +490,34 @@ def _quant_label(value: str) -> str: default_value=None, validation_rules=[ValidationRule(ValidationRuleType.MIN, value=1, message="Interval must be at least 1")], dependencies=[FieldDependency(field="gradient_checkpointing", operator="equals", value=True, action="enable")], - help_text="Checkpoint every N transformer blocks (leave blank to disable)", - tooltip="Higher values save more memory but increase compute. Clear this field to turn off partial checkpointing (supported for Flux, Sana, SD3, Chroma, AuraFlow, and HunyuanVideo).", + help_text="Adjust checkpoint spacing or chunk size for supported transformer blocks (leave blank for per-block checkpointing)", + tooltip="Flux, Flux.2, Krea 2, LTXVideo2, MageFlow, Z-Image, and Wan whole-block paths use contiguous chunks of N blocks. Other families may checkpoint every N-th block. Higher values can reduce recompute but keep more activations in VRAM.", importance=ImportanceLevel.ADVANCED, order=16, documentation="OPTIONS.md#--gradient_checkpointing_interval", ) ) + # Gradient Checkpointing Segment Stride + registry._add_field( + ConfigField( + name="gradient_checkpointing_segment_stride", + arg_name="--gradient_checkpointing_segment_stride", + ui_label="Gradient Checkpointing Segment Stride", + field_type=FieldType.NUMBER, + tab="model", + section="memory_optimization", + default_value=None, + validation_rules=[ValidationRule(ValidationRuleType.MIN, value=1, message="Stride must be at least 1")], + dependencies=[FieldDependency(field="gradient_checkpointing", operator="equals", value=True, action="enable")], + help_text="Start a checkpointed segment every N blocks on supported segmented paths", + tooltip="Use with interval as segment width. interval=2 and stride=4 checkpoints two blocks, runs the next two blocks normally, and repeats on supported whole-block paths.", + importance=ImportanceLevel.ADVANCED, + order=17, + documentation="OPTIONS.md#--gradient_checkpointing_segment_stride", + ) + ) + # Offload During Startup registry._add_field( ConfigField( @@ -511,7 +531,7 @@ def _quant_label(value: str) -> str: help_text="Offload text encoders to CPU during VAE caching", tooltip="Useful for large models that OOM during startup. May significantly increase startup time.", importance=ImportanceLevel.ADVANCED, - order=17, + order=18, documentation="OPTIONS.md#--offload_during_startup", ) ) @@ -1131,6 +1151,26 @@ def _quant_label(value: str) -> str: ) ) + # Qwen Text Encoder Model Path + registry._add_field( + ConfigField( + name="qwen_text_encoder_model_name_or_path", + arg_name="--qwen_text_encoder_model_name_or_path", + ui_label="Qwen Model Path", + field_type=FieldType.TEXT, + tab="model", + section="model_config", + subsection="advanced_paths", + default_value=None, + placeholder="path/to/qwen", + help_text="Path to pretrained Qwen text encoder model", + tooltip="HuggingFace model ID or local path for the Qwen text encoder component.", + importance=ImportanceLevel.ADVANCED, + order=30, + documentation="OPTIONS.md#--qwen_text_encoder_model_name_or_path", + ) + ) + # Revision registry._add_field( ConfigField( @@ -1146,7 +1186,7 @@ def _quant_label(value: str) -> str: help_text="Git branch/tag/commit for model version", tooltip="Specific version of the model to load from HuggingFace. Useful for reproducible training.", importance=ImportanceLevel.ADVANCED, - order=30, + order=31, ) ) @@ -1165,7 +1205,7 @@ def _quant_label(value: str) -> str: help_text="Model variant (e.g., fp16, bf16)", tooltip="Specific variant of the model to load, such as precision variants.", importance=ImportanceLevel.ADVANCED, - order=31, + order=32, ) ) @@ -1187,7 +1227,7 @@ def _quant_label(value: str) -> str: help_text="Default precision for quantized base model weights", tooltip="Precision for non-quantized weights in quantized models. BF16 recommended for stability.", importance=ImportanceLevel.ADVANCED, - order=32, + order=33, ) ) @@ -1208,7 +1248,7 @@ def _quant_label(value: str) -> str: tooltip="Experimental feature for memory savings. May impact training quality. Only available for UNet-based architectures.", importance=ImportanceLevel.EXPERIMENTAL, model_specific=["sd15", "sd20", "sdxl", "deepfloyd"], - order=33, + order=34, ) ) diff --git a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/training.py b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/training.py index 97334d492..780589f47 100644 --- a/simpletuner/simpletuner_sdk/server/services/field_registry/sections/training.py +++ b/simpletuner/simpletuner_sdk/server/services/field_registry/sections/training.py @@ -258,6 +258,68 @@ def register_training_fields(registry: "FieldRegistry") -> None: ) ) + registry._add_field( + ConfigField( + name="gradient_checkpointing_offload_attention", + arg_name="--gradient_checkpointing_offload_attention", + ui_label="Offload Attention Activations", + field_type=FieldType.CHECKBOX, + tab="model", + section="memory_optimization", + default_value=False, + help_text="Move attention-side saved activations to CPU instead of keeping them in VRAM", + tooltip="Works on supported attention/FFN split blocks. Best paired with FFN-only checkpointing when attention transfer is cheaper than recompute.", + importance=ImportanceLevel.ADVANCED, + order=3, + documentation="OPTIONS.md#--gradient_checkpointing_offload_attention", + ) + ) + + registry._add_field( + ConfigField( + name="gradient_checkpointing_offload_pin_memory_max_buckets", + arg_name="--gradient_checkpointing_offload_pin_memory_max_buckets", + ui_label="Offload Pinned Memory Buckets", + field_type=FieldType.NUMBER, + tab="model", + section="memory_optimization", + default_value=12, + validation_rules=[ValidationRule(ValidationRuleType.MIN, value=0, message="Bucket count must be non-negative")], + dependencies=[ + FieldDependency( + field="gradient_checkpointing_offload_attention", operator="equals", value=True, action="show" + ) + ], + help_text="Maximum number of distinct pinned CPU tensor buckets for activation offload", + tooltip="Set 0 to disable pinned-memory pooling. Once the cap is reached, unseen tensor shapes use normal CPU memory.", + importance=ImportanceLevel.ADVANCED, + order=4, + documentation="OPTIONS.md#--gradient_checkpointing_offload_pin_memory_max_buckets", + ) + ) + + registry._add_field( + ConfigField( + name="gradient_checkpointing_offload_prefetch", + arg_name="--gradient_checkpointing_offload_prefetch", + ui_label="Prefetch Offloaded Activations", + field_type=FieldType.CHECKBOX, + tab="model", + section="memory_optimization", + default_value=False, + dependencies=[ + FieldDependency( + field="gradient_checkpointing_offload_attention", operator="equals", value=True, action="show" + ) + ], + help_text="Learn activation restore order and prefetch CPU-offloaded tensors back to GPU", + tooltip="Experimental. Can overlap H2D restore with backward compute after the runtime has observed a stable unpack order.", + importance=ImportanceLevel.ADVANCED, + order=5, + documentation="OPTIONS.md#--gradient_checkpointing_offload_prefetch", + ) + ) + # Group Offloading registry._add_field( ConfigField( diff --git a/simpletuner/static/js/utils/api.js b/simpletuner/static/js/utils/api.js index 210a4e020..6286a5d7b 100644 --- a/simpletuner/static/js/utils/api.js +++ b/simpletuner/static/js/utils/api.js @@ -12,18 +12,22 @@ return path.startsWith("/") ? path : `/${path}`; } + function normalizeBaseUrl(url) { + return typeof url === "string" ? url.replace(/\/+$/, "") : ""; + } + function getBaseOrigin() { return global.location ? global.location.origin : ""; } const ApiClient = { get apiBaseUrl() { - return (global.ServerConfig && global.ServerConfig.apiBaseUrl) || getBaseOrigin(); + return normalizeBaseUrl((global.ServerConfig && global.ServerConfig.apiBaseUrl) || getBaseOrigin()); }, get callbackBaseUrl() { if (global.ServerConfig && global.ServerConfig.callbackUrl) { - return global.ServerConfig.callbackUrl; + return normalizeBaseUrl(global.ServerConfig.callbackUrl); } return this.apiBaseUrl; }, diff --git a/tests/helpers/metadata/test_audio_backends.py b/tests/helpers/metadata/test_audio_backends.py index 25754235c..161fb5f8d 100644 --- a/tests/helpers/metadata/test_audio_backends.py +++ b/tests/helpers/metadata/test_audio_backends.py @@ -218,6 +218,45 @@ def test_parquet_audio_bucket_uses_manifest_metadata(self): self.assertEqual(audio_metadata["num_channels"], 2) self.assertEqual(audio_metadata["truncation_mode"], "beginning") + def test_parquet_audio_bucket_preserves_lyrics_and_token_path(self): + dataset_id = "audio_parquet_tokens" + audio_path = os.path.join(self.tempdir.name, "tokenized.wav") + sample_rate, num_samples = _write_test_wav(audio_path) + + manifest_path = os.path.join(self.tempdir.name, "tokens.jsonl") + record = { + "filepath": os.path.basename(audio_path), + "description": "bright synth pop", + "sample_rate": sample_rate, + "num_samples": num_samples, + "duration": round(num_samples / sample_rate, 2), + "channels": 1, + "lyrics": "hello from heartmula", + "audio_tokens_path": "tokenized.tokens.npy", + } + with open(manifest_path, "w", encoding="utf-8") as manifest_file: + manifest_file.write(json.dumps(record) + "\n") + + backend = self._build_backend(dataset_id) + self._register_dataset_config(dataset_id) + parquet_backend = self._parquet_backend(dataset_id, backend, manifest_path) + parquet_backend.parquet_config["caption_column"] = "description" + parquet_backend.parquet_config["lyrics_column"] = "lyrics" + + metadata_updates = {} + parquet_backend._process_for_bucket( + image_path_str=audio_path, + aspect_ratio_bucket_indices={}, + metadata_updates=metadata_updates, + statistics={}, + ) + + audio_metadata = metadata_updates[audio_path] + self.assertEqual(audio_metadata["tags"], "bright synth pop") + self.assertEqual(audio_metadata["prompt"], "bright synth pop") + self.assertEqual(audio_metadata["lyrics"], "hello from heartmula") + self.assertEqual(audio_metadata["audio_tokens_path"], "tokenized.tokens.npy") + def test_parquet_audio_bucket_reads_from_audio_when_manifest_missing_values(self): dataset_id = "audio_parquet_fallback" audio_path = os.path.join(self.tempdir.name, "fallback.wav") diff --git a/tests/js/api_client.test.js b/tests/js/api_client.test.js index 13021ffe6..940d48982 100644 --- a/tests/js/api_client.test.js +++ b/tests/js/api_client.test.js @@ -84,6 +84,22 @@ describe('ApiClient', () => { }); }); + describe('base URL normalization', () => { + test('removes trailing slashes from configured base URLs', () => { + global.ServerConfig = { + apiBaseUrl: 'https://api.example.com/', + callbackUrl: 'https://callback.example.com///', + }; + + expect(window.ApiClient.apiBaseUrl).toBe('https://api.example.com'); + expect(window.ApiClient.callbackBaseUrl).toBe('https://callback.example.com'); + expect(window.ApiClient.resolve('/api/jobs', { forceApi: true })) + .toBe('https://api.example.com/api/jobs'); + expect(window.ApiClient.resolve('/notify', { forceCallback: true })) + .toBe('https://callback.example.com/notify'); + }); + }); + describe('resolve', () => { test('resolves path with default options', () => { const result = window.ApiClient.resolve('/api/jobs'); diff --git a/tests/test_ace_step_model.py b/tests/test_ace_step_model.py index 4ccbbf89e..a96b807ec 100644 --- a/tests/test_ace_step_model.py +++ b/tests/test_ace_step_model.py @@ -24,6 +24,7 @@ def setUp(self): self.config.pretrained_model_name_or_path = "dummy_path" self.config.pretrained_transformer_model_name_or_path = None self.config.pretrained_transformer_subfolder = None + self.config.qwen_text_encoder_model_name_or_path = None self.config.model_flavour = "base" self.config.controlnet = False self.config.peft_lora_target_modules = None @@ -185,6 +186,21 @@ def test_resolve_v15_layout_uses_requested_variant(self): self.assertEqual(layout["tokenizer_path"], str(root / "Qwen3-Embedding-0.6B")) self.assertEqual(layout["vae_path"], str(root / "vae")) + def test_resolve_v15_layout_uses_qwen_text_encoder_override(self): + self.config.model_flavour = "v15-base" + self.config.qwen_text_encoder_model_name_or_path = "custom/qwen" + with TemporaryDirectory() as tmpdir: + root = Path(tmpdir) + (root / "vae").mkdir() + (root / "acestep-v15-base").mkdir() + torch.save(torch.zeros(1, 64, 4), root / "acestep-v15-base" / "silence_latent.pt") + + layout = self.model._resolve_v15_layout(str(root)) + + self.assertIsNotNone(layout) + self.assertEqual(layout["tokenizer_path"], "custom/qwen") + self.assertEqual(layout["vae_path"], str(root / "vae")) + def test_resolve_v15_layout_caches_negative_result_for_same_base_path(self): with TemporaryDirectory() as tmpdir: root = Path(tmpdir) diff --git a/tests/test_acestep_lora_targets.py b/tests/test_acestep_lora_targets.py index 3edb8171b..ad547b438 100644 --- a/tests/test_acestep_lora_targets.py +++ b/tests/test_acestep_lora_targets.py @@ -20,6 +20,7 @@ def setUp(self): self.config.pretrained_model_name_or_path = "dummy_path" self.config.pretrained_transformer_model_name_or_path = None self.config.pretrained_transformer_subfolder = None + self.config.qwen_text_encoder_model_name_or_path = None self.config.model_family = "ace_step" self.config.model_flavour = "base" self.config.peft_lora_target_modules = None diff --git a/tests/test_attention_backend.py b/tests/test_attention_backend.py index 7a65005af..0f40948e6 100644 --- a/tests/test_attention_backend.py +++ b/tests/test_attention_backend.py @@ -1,8 +1,9 @@ import os import shutil +import sys import tempfile import unittest -from types import SimpleNamespace +from types import ModuleType, SimpleNamespace from unittest.mock import patch import torch @@ -777,3 +778,54 @@ def test_packed_backend_accepts_config_fa2_hub_varlen_alias(self): attention_backend_module._select_packed_backend("flash-attn-varlen-hub"), "flash-attn-varlen-hub", ) + + def test_packed_hub_backend_passes_trust_remote_code_without_touching_telemetry(self): + calls = [] + constants_module = ModuleType("huggingface_hub.constants") + constants_module.HF_HUB_DISABLE_TELEMETRY = True + kernels_module = ModuleType("kernels") + + class DummyKernel: + @staticmethod + def flash_attn_qkvpacked_func(qkv, dropout_p=0.0, softmax_scale=None, causal=False): + return qkv[:, :, 0] + + def fake_get_kernel(target, **kwargs): + calls.append((target, dict(kwargs), constants_module.HF_HUB_DISABLE_TELEMETRY)) + return DummyKernel() + + kernels_module.get_kernel = fake_get_kernel + + with patch.dict( + sys.modules, + { + "huggingface_hub.constants": constants_module, + "kernels": kernels_module, + }, + ): + backend = attention_backend_module.get_packed_attention_backend("flash-attn-varlen-hub") + + self.assertEqual(backend.name, "flash-attn-varlen-hub") + self.assertEqual(calls, [("kernels-community/flash-attn2", {"trust_remote_code": True, "version": 3}, True)]) + self.assertTrue(constants_module.HF_HUB_DISABLE_TELEMETRY) + + def test_packed_hub_backend_omits_unknown_kernel_version(self): + calls = [] + kernels_module = ModuleType("kernels") + + class DummyKernel: + @staticmethod + def flash_attn_qkvpacked_func(qkv, dropout_p=0.0, softmax_scale=None, causal=False): + return qkv[:, :, 0] + + def fake_get_kernel(target, **kwargs): + calls.append((target, dict(kwargs))) + return DummyKernel() + + kernels_module.get_kernel = fake_get_kernel + + with patch.dict(sys.modules, {"kernels": kernels_module}): + kernel = attention_backend_module._get_hub_kernel("example/custom-kernel") + + self.assertIsInstance(kernel, DummyKernel) + self.assertEqual(calls, [("example/custom-kernel", {"trust_remote_code": True})]) diff --git a/tests/test_audio_sampler.py b/tests/test_audio_sampler.py index 0115ef362..41a9c01d0 100644 --- a/tests/test_audio_sampler.py +++ b/tests/test_audio_sampler.py @@ -19,9 +19,11 @@ def setUp(self): self.mock_metadata.id = "test_backend" self.mock_metadata.aspect_ratio_bucket_indices = {"10s": ["a.wav", "b.wav"], "20s": ["c.wav"]} self.mock_metadata.seen_images = {} - self.mock_metadata.is_seen = lambda x: x in self.mock_metadata.seen_images + self.mock_metadata.is_seen = ( + lambda x, occurrence_index=0: int(self.mock_metadata.seen_images.get(x, 0)) > occurrence_index + ) self.mock_metadata.mark_batch_as_seen = lambda batch: [ - self.mock_metadata.seen_images.update({x: True}) for x in batch + self.mock_metadata.seen_images.update({x: int(self.mock_metadata.seen_images.get(x, 0)) + 1}) for x in batch ] self.mock_metadata.reset_seen_images = lambda: self.mock_metadata.seen_images.clear() self.mock_metadata.instance_data_dir = "" diff --git a/tests/test_boogu_image_model.py b/tests/test_boogu_image_model.py index 6481ec054..7d243b8c2 100644 --- a/tests/test_boogu_image_model.py +++ b/tests/test_boogu_image_model.py @@ -10,6 +10,7 @@ from simpletuner.helpers.models.boogu_image.model import BooguImage from simpletuner.helpers.models.boogu_image.pipeline import BooguImagePipeline, retrieve_timesteps from simpletuner.helpers.models.boogu_image.rope import BooguImageDoubleStreamRotaryPosEmbed +from simpletuner.helpers.models.boogu_image.transformer import BooguImageTransformer2DModel from simpletuner.helpers.models.common import PipelineTypes from simpletuner.helpers.models.flux.model import Flux from simpletuner.helpers.training.attention_backend import _DIFFUSERS_BACKEND_ALIASES @@ -187,6 +188,63 @@ def test_torchao_filter_skips_boogu_reference_image_modules(self): self.assertFalse(_torchao_filter_fn(torch.nn.Linear(16, 16), "ref_image_refiner.0.attn.to_q")) self.assertTrue(_torchao_filter_fn(torch.nn.Linear(16, 16), "context_refiner.0.attn.to_q")) + def test_patch_embed_refine_clones_custom_autograd_outputs_before_layout_writes(self): + class PassThroughFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, tensor): + return tensor + + @staticmethod + def backward(ctx, grad_output): + return grad_output + + class PassThroughModule(torch.nn.Module): + def forward(self, tensor): + return PassThroughFunction.apply(tensor) + + transformer = BooguImageTransformer2DModel( + patch_size=2, + in_channels=1, + hidden_size=4, + num_layers=0, + num_double_stream_layers=0, + num_refiner_layers=0, + num_attention_heads=1, + num_kv_heads=1, + axes_dim_rope=(1, 1, 2), + axes_lens=(4, 4, 4), + instruction_feature_configs={ + "instruction_feat_dim": 4, + "reduce_type": "mean", + "num_instruction_feat_layers": 1, + }, + ) + transformer.x_embedder = PassThroughModule() + transformer.ref_image_patch_embedder = PassThroughModule() + + hidden_states = torch.randn(1, 4, 4, requires_grad=True) + ref_image_hidden_states = torch.randn(1, 2, 4, requires_grad=True) + padded_img_mask = torch.ones(1, 4, dtype=torch.bool) + padded_ref_img_mask = torch.ones(1, 2, dtype=torch.bool) + noise_rotary_emb = torch.zeros(1, 4, 4) + ref_img_rotary_emb = torch.zeros(1, 2, 4) + temb = torch.zeros(1, 4) + + output = transformer.img_patch_embed_and_refine( + hidden_states, + ref_image_hidden_states, + padded_img_mask, + padded_ref_img_mask, + noise_rotary_emb, + ref_img_rotary_emb, + [[2]], + [4], + temb, + ) + + self.assertEqual(tuple(output.shape), (1, 6, 4)) + output.sum().backward() + def test_boogu_rotary_complex_inputs_use_real_valued_math(self): x = torch.randn(2, 5, 3, 8, dtype=torch.bfloat16) angles = torch.randn(2, 5, 4) diff --git a/tests/test_collate.py b/tests/test_collate.py index bcd1a8515..1edee112b 100644 --- a/tests/test_collate.py +++ b/tests/test_collate.py @@ -26,12 +26,14 @@ def __init__( requires_conditioning_latents: bool = False, requires_conditioning_dataset: bool = False, requires_text_embed_image_context: bool = False, + uses_text_embeddings_cache: bool = True, ): self._requires_conditioning = requires_conditioning self._use_reference_embeds = use_reference_embeds self._requires_conditioning_latents = requires_conditioning_latents self._requires_conditioning_dataset = requires_conditioning_dataset self._requires_text_embed_context = requires_text_embed_image_context + self._uses_text_embeddings_cache = uses_text_embeddings_cache def requires_conditioning_image_embeds(self): return self._requires_conditioning @@ -54,6 +56,9 @@ def get_transforms(self, dataset_type: str | None = None): def requires_text_embed_image_context(self): return self._requires_text_embed_context + def uses_text_embeddings_cache(self): + return self._uses_text_embeddings_cache + class _StubConditioningSample: def __init__( @@ -180,6 +185,25 @@ def test_collate_fn_basic_path_uses_backend_for_text_cache(self): backend_mock = active_mocks[5] backend_mock.assert_called() + def test_collate_fn_preserves_prompts_without_text_cache(self): + backend_dict = { + "data_backend": _make_stub_data_backend(), + "config": {"instance_data_dir": "/train"}, + } + + model = _StubModel(requires_conditioning=False, uses_text_embeddings_cache=False) + patchers, _ = self._patch_state_tracker( + model=model, data_backend=backend_dict, text_outputs={}, backend_lookup={"backend-1": backend_dict} + ) + with ExitStack() as stack: + active_mocks = [stack.enter_context(patcher) for patcher in patchers] + result = collate_fn(self.base_batch) + + self.assertEqual(result["prompts"], ["caption"]) + self.assertEqual(result["text_encoder_output"], {}) + self.assertIsNone(result["prompt_embeds"]) + active_mocks[9].assert_not_called() + def test_compute_latents_uses_backend_ondemand_mode(self): vae_cache = SimpleNamespace( vae_cache_ondemand=True, diff --git a/tests/test_conditioning_split_alignment.py b/tests/test_conditioning_split_alignment.py index e76756990..7c2077d54 100644 --- a/tests/test_conditioning_split_alignment.py +++ b/tests/test_conditioning_split_alignment.py @@ -1,3 +1,4 @@ +import json import multiprocessing import os import unittest @@ -157,6 +158,7 @@ def _prepare_metadata_backends( source_data_backend: InMemoryDataBackend | None = None, conditioning_data_backend: InMemoryDataBackend | None = None, before_copy_callback=None, + apply_padding: bool = False, ): source_config = { "resolution_type": "area", @@ -195,7 +197,7 @@ def _prepare_metadata_backends( "dataset_type": "image", } ) - training_metadata.split_buckets_between_processes() + training_metadata.split_buckets_between_processes(apply_padding=apply_padding) conditioning_config = { "resolution_type": "area", @@ -425,6 +427,106 @@ def before_copy(train_meta, cond_meta, train_backend, cond_backend, source_entry "Conditioning duplication should retain samples even when cache reload returns empty data.", ) + def test_rank_local_conditioning_schedule_does_not_overwrite_target_cache(self): + base_images = {"1.0": [f"/datasets/train/img_{i}.png" for i in range(16)]} + accelerator = self._init_state(num_processes=4, process_index=1, train_batch_size=8) + source_schedule = {} + source_sentinel = {"config": {}, "aspect_ratio_bucket_indices": {"1.0": ["canonical-source-cache"]}} + sentinel = {"config": {}, "aspect_ratio_bucket_indices": {"1.0": ["canonical-target-cache"]}} + + def before_copy(train_meta, cond_meta, train_backend, cond_backend, *_args): + source_schedule.update({bucket: list(paths) for bucket, paths in train_meta.aspect_ratio_bucket_indices.items()}) + train_backend.write(train_meta.cache_file, json.dumps(source_sentinel)) + cond_backend.write(cond_meta.cache_file, json.dumps(sentinel)) + + train_meta, cond_meta = self._prepare_metadata_backends( + accelerator=accelerator, + base_buckets=base_images, + source_id="cache_source", + source_dir="/datasets/train", + conditioning_id="cache_conditioning", + conditioning_dir="/datasets/control", + before_copy_callback=before_copy, + ) + + self.assertTrue(train_meta.read_only) + self.assertTrue(cond_meta.read_only) + self.assertEqual(source_schedule, train_meta.aspect_ratio_bucket_indices) + self.assertEqual(json.loads(train_meta.data_backend.read(train_meta.cache_file)), source_sentinel) + self.assertEqual( + json.loads(cond_meta.data_backend.read(cond_meta.cache_file))["aspect_ratio_bucket_indices"], + sentinel["aspect_ratio_bucket_indices"], + ) + self.assertEqual( + cond_meta.aspect_ratio_bucket_indices["1.0"], + [path.replace("/datasets/train", "/datasets/control", 1) for path in source_schedule["1.0"]], + ) + + def test_duplication_preserves_the_padded_split_source_schedule(self): + # The main process persists the unsplit discovery list before the split, and the + # post-split save is a no-op because the split marks the backend read-only. Duplication + # must not reload that unsplit list back over the rank-local schedule. + source_dir = "/datasets/pad_train" + conditioning_dir = "/datasets/pad_control" + base_images = {"1.0": [f"{source_dir}/img_{index}.png" for index in range(4)]} + num_processes = 8 + conditioning_backend = InMemoryDataBackend("pad_control") + conditioning_cache_path = f"{conditioning_dir}/aspect_ratio_bucket_indices.json" + + source_shards = [] + conditioning_cache_snapshots = [] + + def persist_unsplit_source_cache(train_meta, *_args): + split_view = {bucket: list(paths) for bucket, paths in train_meta.aspect_ratio_bucket_indices.items()} + original_read_only = train_meta.read_only + train_meta.read_only = False + train_meta.aspect_ratio_bucket_indices = deepcopy(base_images) + train_meta.save_cache() + train_meta.aspect_ratio_bucket_indices = split_view + train_meta.read_only = original_read_only + + for process_index in range(num_processes): + with self.subTest(process=process_index): + accelerator = self._init_state(num_processes=num_processes, process_index=process_index) + source_backend = InMemoryDataBackend("pad_train") + train_meta, cond_meta = self._prepare_metadata_backends( + accelerator=accelerator, + base_buckets=base_images, + source_id="pad_train", + source_dir=source_dir, + conditioning_id="pad_control", + conditioning_dir=conditioning_dir, + source_data_backend=source_backend, + conditioning_data_backend=conditioning_backend, + before_copy_callback=persist_unsplit_source_cache, + apply_padding=True, + ) + + # The clobber source has to be present, or this test proves nothing. + self.assertIn(f"{source_dir}/aspect_ratio_bucket_indices.json", source_backend._storage) + + # The clobber also coerces bucket keys to float, so do not index by "1.0". + (shard,) = (list(paths) for paths in train_meta.aspect_ratio_bucket_indices.values()) + source_shards.append(shard) + conditioning_cache_snapshots.append(conditioning_backend._storage.get(conditioning_cache_path)) + + (conditioning_shard,) = (list(paths) for paths in cond_meta.aspect_ratio_bucket_indices.values()) + self.assertEqual( + conditioning_shard, + [path.replace(source_dir, conditioning_dir, 1) for path in shard], + msg=f"Conditioning dataset did not inherit the rank {process_index} shard.", + ) + + self.assertEqual([len(shard) for shard in source_shards], [1] * num_processes) + self.assertEqual( + sorted({path for shard in source_shards for path in shard}), + base_images["1.0"], + msg="The eight rank shards no longer cover the dataset.", + ) + # Every rank writes the conditioning cache to the same path, so its contents must not + # depend on which rank wrote last. + self.assertEqual(len(set(conditioning_cache_snapshots)), 1) + def test_reference_strict_duplication_multi_process_reload(self): manager = multiprocessing.Manager() try: diff --git a/tests/test_cosmos3_model.py b/tests/test_cosmos3_model.py index 8ab3b01a5..3c33fc548 100644 --- a/tests/test_cosmos3_model.py +++ b/tests/test_cosmos3_model.py @@ -4,6 +4,7 @@ from importlib import resources from pathlib import Path from types import SimpleNamespace +from unittest import mock import torch from safetensors import safe_open @@ -252,6 +253,55 @@ def test_freeze_reasoning_layers_keeps_generation_path_trainable(self): self.assertTrue(transformer.layers[0].self_attn.add_q_proj.weight.requires_grad) self.assertTrue(transformer.layers[0].mlp_moe_gen.down_proj.weight.requires_grad) + def test_unpatchify_preserves_prediction_dtype(self): + transformer = Cosmos3OmniTransformer( + hidden_size=8, + intermediate_size=16, + head_dim=4, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=1, + latent_channel=2, + patch_latent_dim=8, + vocab_size=32, + ) + packed = torch.randn(4, 8, dtype=torch.bfloat16) + + unpacked = transformer._unpatchify_and_unpack_latents( + packed_mse_preds=packed, + token_shapes_vision=[(1, 2, 2)], + noisy_frame_indexes_vision=[torch.tensor([0], dtype=torch.long)], + original_latent_shapes=[(1, 4, 4)], + ) + + self.assertEqual(unpacked[0].dtype, torch.bfloat16) + self.assertEqual(unpacked[0].shape, (1, 2, 1, 4, 4)) + + def test_unpatchify_uses_materialized_latent_dtype(self): + transformer = Cosmos3OmniTransformer( + hidden_size=8, + intermediate_size=16, + head_dim=4, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=1, + latent_channel=2, + patch_latent_dim=8, + vocab_size=32, + ) + packed = torch.randn(4, 8, dtype=torch.float32) + + with mock.patch("torch.einsum", return_value=torch.zeros(2, 1, 2, 2, 2, 2, dtype=torch.bfloat16)): + unpacked = transformer._unpatchify_and_unpack_latents( + packed_mse_preds=packed, + token_shapes_vision=[(1, 2, 2)], + noisy_frame_indexes_vision=[torch.tensor([0], dtype=torch.long)], + original_latent_shapes=[(1, 4, 4)], + ) + + self.assertEqual(unpacked[0].dtype, torch.bfloat16) + self.assertEqual(unpacked[0].shape, (1, 2, 1, 4, 4)) + def test_cosmos3_text_cache_metadata_and_collation(self): model = TestableCosmos3Image() latent = torch.zeros(16, 3, 4, 5) diff --git a/tests/test_factory_edge_cases.py b/tests/test_factory_edge_cases.py index 98febb561..0aacb9628 100644 --- a/tests/test_factory_edge_cases.py +++ b/tests/test_factory_edge_cases.py @@ -977,6 +977,7 @@ def test_early_validation_impossible_config(self): # Mock accelerator with 8 GPUs mock_accelerator = MagicMock() mock_accelerator.num_processes = 8 + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator # Mock StateTracker.get_args() to return args with allow_dataset_oversubscription=False @@ -1023,16 +1024,7 @@ def test_early_validation_sufficient_repeats(self): # Mock accelerator with 2 GPUs mock_accelerator = MagicMock() mock_accelerator.num_processes = 2 - - # Mock split_between_processes to distribute images properly - def split_side_effect(images, apply_padding=False): - mock_context = MagicMock() - # Each GPU gets half the images - mock_context.__enter__ = MagicMock(return_value=images[: len(images) // 2]) - mock_context.__exit__ = MagicMock(return_value=False) - return mock_context - - mock_accelerator.split_between_processes = MagicMock(side_effect=split_side_effect) + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator # Mock StateTracker.get_args() to return args with allow_dataset_oversubscription=False @@ -1068,14 +1060,7 @@ def test_eval_dataset_ignores_gradient_accumulation(self): mock_accelerator = MagicMock() mock_accelerator.num_processes = 1 - - def split_side_effect(images, apply_padding=False): - mock_context = MagicMock() - mock_context.__enter__ = MagicMock(return_value=images) - mock_context.__exit__ = MagicMock(return_value=False) - return mock_context - - mock_accelerator.split_between_processes = MagicMock(side_effect=split_side_effect) + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator with ( @@ -1092,8 +1077,6 @@ def split_side_effect(images, apply_padding=False): if "Dataset configuration will produce zero usable batches" in str(e): self.fail(f"Eval dataset should not consider grad accumulation: {e}") - mock_accelerator.split_between_processes.assert_called_once() - def test_oversubscription_auto_adjustment(self): """Test that --allow_dataset_oversubscription automatically adjusts repeats.""" from simpletuner.helpers.metadata.backends.base import MetadataBackend @@ -1110,6 +1093,7 @@ def test_oversubscription_auto_adjustment(self): # Mock accelerator with 8 GPUs mock_accelerator = MagicMock() mock_accelerator.num_processes = 8 + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator # Mock StateTracker to enable oversubscription and NO user-set repeats @@ -1122,20 +1106,11 @@ def test_oversubscription_auto_adjustment(self): # Return config WITHOUT 'repeats' key (not user-set) mock_get_config.return_value = {"id": "test_auto_adjust"} - # Mock the split_between_processes to avoid actual splitting - def split_side_effect(images, apply_padding=False): - mock_context = MagicMock() - mock_context.__enter__ = MagicMock(return_value=images) - mock_context.__exit__ = MagicMock(return_value=False) - return mock_context - - mock_accelerator.split_between_processes = MagicMock(side_effect=split_side_effect) - # Should NOT raise, should auto-adjust repeats try: MetadataBackend.split_buckets_between_processes(mock_backend, gradient_accumulation_steps=1) - # Check that repeats was adjusted - self.assertEqual(mock_backend.repeats, 7, "Repeats should be auto-adjusted to 7") + self.assertEqual(mock_backend.repeats, 0) + self.assertEqual(len(mock_backend.aspect_ratio_bucket_indices["1.0"]), 4) except ValueError as e: self.fail(f"Should not raise with oversubscription enabled: {e}") @@ -1154,6 +1129,7 @@ def test_oversubscription_respects_manual_repeats(self): mock_accelerator = MagicMock() mock_accelerator.num_processes = 8 + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator # Mock StateTracker with oversubscription enabled BUT user set repeats @@ -1189,6 +1165,7 @@ def test_oversubscription_disabled_raises_error(self): mock_accelerator = MagicMock() mock_accelerator.num_processes = 8 + mock_accelerator.process_index = 0 mock_backend.accelerator = mock_accelerator # Mock StateTracker with oversubscription DISABLED @@ -1208,6 +1185,49 @@ def test_oversubscription_disabled_raises_error(self): self.assertIn("Dataset configuration will produce zero usable batches", error_msg) self.assertIn("Enable --allow_dataset_oversubscription", error_msg) + def test_bucket_split_padding_policy_matches_training_mode(self): + """Factory padding covers epoch and explicit oversubscription policies.""" + from simpletuner.helpers.data_backend.factory import FactoryRegistry + + factory = FactoryRegistry( + args=self.args, + accelerator=self.accelerator, + text_encoders=self.text_encoders, + tokenizers=self.tokenizers, + model=self.model, + ) + factory._handle_config_versioning = MagicMock() + + cases = ( + ("fixed_steps", 100, False, False), + ("oversubscribed_fixed_steps", 100, True, True), + ("epoch_driven", 0, False, True), + ) + for name, max_train_steps, allow_oversubscription, expected in cases: + with self.subTest(name=name): + factory.args.max_train_steps = max_train_steps + factory.args.allow_dataset_oversubscription = allow_oversubscription + factory.args.skip_file_discovery = "aspect" + factory.args.eval_dataset_id = None + + metadata_backend = MagicMock() + init_backend = { + "id": "train", + "config": {}, + "dataset_type": "image", + "metadata_backend": metadata_backend, + } + factory._handle_bucket_operations( + backend={"id": "train", "skip_file_discovery": "aspect"}, + init_backend=init_backend, + conditioning_type=None, + ) + + metadata_backend.split_buckets_between_processes.assert_called_once_with( + gradient_accumulation_steps=factory.args.gradient_accumulation_steps, + apply_padding=expected, + ) + def test_image_embeds_backend_configuration(self): """Image embed configuration should not instantiate VAE cache directly.""" from simpletuner.helpers.data_backend.factory import FactoryRegistry diff --git a/tests/test_gradient_checkpointing_backend.py b/tests/test_gradient_checkpointing_backend.py index 54d68308e..527e2381c 100644 --- a/tests/test_gradient_checkpointing_backend.py +++ b/tests/test_gradient_checkpointing_backend.py @@ -3,6 +3,8 @@ """ import unittest +from types import SimpleNamespace +from unittest import mock import torch import torch.nn as nn @@ -109,12 +111,12 @@ def test_cpu_offload_hooks_pack_unpack(self): hooks = CPUOffloadHooks() if torch.cuda.is_available(): - tensor = torch.randn(4, 4, device="cuda") + tensor = torch.randn(4, 4, device="cuda", requires_grad=True) * 2 packed = hooks.pack(tensor) - # Pack returns (cpu_tensor, original_device) tuple + # Pack returns (cpu_tensor, original_device[, pool_key]) tuple self.assertIsInstance(packed, tuple) - self.assertEqual(len(packed), 2) - cpu_tensor, original_device = packed + self.assertGreaterEqual(len(packed), 2) + cpu_tensor, original_device = packed[:2] self.assertEqual(cpu_tensor.device.type, "cpu") self.assertEqual(original_device.type, "cuda") @@ -129,10 +131,504 @@ def test_cpu_offload_hooks_pack_unpack(self): self.assertEqual(cpu_tensor.device.type, "cpu") self.assertIsNone(original_device) + def test_activation_offload_does_not_recompute_forward(self): + """Activation offload preserves the original forward graph instead of rematerializing it.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import activation_offload + + class CountingModule(nn.Module): + def __init__(self): + super().__init__() + self.calls = 0 + self.linear = nn.Linear(4, 4) + + def forward(self, x): + self.calls += 1 + return torch.relu(self.linear(x)) + + module = CountingModule() + x = torch.randn(2, 4, requires_grad=True) + activation_offload(module, x).sum().backward() + + self.assertEqual(module.calls, 1) + + def test_activation_offload_pin_memory_bucket_setting(self): + """Pinned bucket limit is configurable.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + get_activation_offload_pin_memory_max_buckets, + set_activation_offload_pin_memory_max_buckets, + ) + + original = get_activation_offload_pin_memory_max_buckets() + try: + set_activation_offload_pin_memory_max_buckets(7) + self.assertEqual(get_activation_offload_pin_memory_max_buckets(), 7) + set_activation_offload_pin_memory_max_buckets(0) + self.assertEqual(get_activation_offload_pin_memory_max_buckets(), 0) + finally: + set_activation_offload_pin_memory_max_buckets(original) + + def test_activation_offload_pin_memory_bucket_count_normalization(self): + """Pinned bucket config values produce explicit validation errors.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + normalize_activation_offload_pin_memory_max_buckets, + ) + + self.assertEqual(normalize_activation_offload_pin_memory_max_buckets(None), 12) + self.assertEqual(normalize_activation_offload_pin_memory_max_buckets(""), 12) + self.assertEqual(normalize_activation_offload_pin_memory_max_buckets("0"), 0) + self.assertEqual(normalize_activation_offload_pin_memory_max_buckets(3), 3) + + with self.assertRaisesRegex(ValueError, "non-negative integer"): + normalize_activation_offload_pin_memory_max_buckets("many") + with self.assertRaisesRegex(ValueError, "non-negative"): + normalize_activation_offload_pin_memory_max_buckets(-1) + + def test_activation_offload_copy_stream_counts_are_configurable(self): + """Activation offload copy stream pools expose bounded tunable widths.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + get_activation_offload_copy_stream_stats, + get_activation_offload_d2h_copy_stream_count, + get_activation_offload_h2d_prefetch_stream_count, + set_activation_offload_d2h_copy_stream_count, + set_activation_offload_h2d_prefetch_stream_count, + ) + + original_d2h = get_activation_offload_d2h_copy_stream_count() + original_h2d = get_activation_offload_h2d_prefetch_stream_count() + try: + set_activation_offload_d2h_copy_stream_count(3) + set_activation_offload_h2d_prefetch_stream_count(5) + + self.assertEqual(get_activation_offload_d2h_copy_stream_count(), 3) + self.assertEqual(get_activation_offload_h2d_prefetch_stream_count(), 5) + stats = get_activation_offload_copy_stream_stats() + self.assertEqual(stats["d2h"]["width"], 3) + self.assertEqual(stats["h2d_prefetch"]["width"], 5) + + set_activation_offload_d2h_copy_stream_count(0) + set_activation_offload_h2d_prefetch_stream_count(-1) + self.assertEqual(get_activation_offload_d2h_copy_stream_count(), 1) + self.assertEqual(get_activation_offload_h2d_prefetch_stream_count(), 1) + finally: + set_activation_offload_d2h_copy_stream_count(original_d2h) + set_activation_offload_h2d_prefetch_stream_count(original_h2d) + + def test_pinned_memory_pool_tracks_reuse_stats(self): + """Pinned pool stats track allocation, release, and buffer reuse per bucket.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import _PinnedMemoryPool + + pool = _PinnedMemoryPool(max_buckets=1) + pool._allocate = lambda key: torch.empty_strided(key.size, key.stride, dtype=key.dtype, layout=key.layout) + key = pool.key_for(torch.empty(2, 3)) + + first = pool.checkout(key) + pool.release_after_cuda_copy(key, first, torch.device("cpu")) + second = pool.checkout(key) + + self.assertIs(first, second) + snapshot = pool.snapshot() + bucket = snapshot["buckets"][0] + self.assertEqual(snapshot["total_accesses"], 2) + self.assertEqual(snapshot["total_allocations"], 1) + self.assertEqual(snapshot["total_buffer_reuses"], 1) + self.assertEqual(bucket["accesses"], 2) + self.assertEqual(bucket["allocations"], 1) + self.assertEqual(bucket["buffer_reuses"], 1) + self.assertEqual(bucket["releases"], 1) + + def test_pinned_memory_pool_eviction_uses_persisted_stats(self): + """Repeated non-resident shapes can evict colder resident buckets without losing old counters.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import _PinnedMemoryPool + + pool = _PinnedMemoryPool(max_buckets=1) + pool._allocate = lambda key: torch.empty_strided(key.size, key.stride, dtype=key.dtype, layout=key.layout) + key_a = pool.key_for(torch.empty(2, 3)) + key_b = pool.key_for(torch.empty(4, 5)) + + first = pool.checkout(key_a) + pool.release_after_cuda_copy(key_a, first, torch.device("cpu")) + + self.assertIsNone(pool.checkout(key_b)) + second = pool.checkout(key_b) + + self.assertIsNotNone(second) + snapshot = pool.snapshot() + buckets = {(bucket["size"], bucket["stride"]): bucket for bucket in snapshot["buckets"]} + bucket_a = buckets[((2, 3), (3, 1))] + bucket_b = buckets[((4, 5), (5, 1))] + + self.assertFalse(bucket_a["resident"]) + self.assertEqual(bucket_a["evictions"], 1) + self.assertEqual(bucket_a["accesses"], 1) + self.assertTrue(bucket_b["resident"]) + self.assertEqual(bucket_b["accesses"], 2) + self.assertEqual(bucket_b["cap_misses"], 1) + self.assertEqual(bucket_b["admissions"], 1) + self.assertEqual(snapshot["tracked_buckets"], 2) + self.assertEqual(snapshot["total_evictions"], 1) + + def test_cpu_offload_dense_noncontiguous_views_use_flat_transfer(self): + """Dense views can transfer as flat storage and restore their original logical stride.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import CPUOffloadHooks + + hooks = CPUOffloadHooks() + base = torch.arange(12).view(3, 4) + transposed = base.t() + + flat = hooks._flat_storage_view(transposed) + + self.assertIsNotNone(flat) + self.assertEqual(tuple(flat.shape), (12,)) + self.assertEqual(tuple(flat.stride()), (1,)) + self.assertEqual(flat.storage_offset(), transposed.storage_offset()) + + transfer_tensor, restore_view = hooks._transfer_view(transposed) + restored = hooks.unpack((transfer_tensor.clone(), torch.device("cpu"), None, restore_view)) + + self.assertEqual(tuple(transfer_tensor.shape), (12,)) + self.assertEqual(tuple(restored.shape), tuple(transposed.shape)) + self.assertEqual(tuple(restored.stride()), tuple(transposed.stride())) + self.assertTrue(torch.equal(restored, transposed)) + + def test_cpu_offload_dense_offset_views_restore_values(self): + """Dense views with non-zero base offsets restore from the copied transfer span.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import CPUOffloadHooks + + hooks = CPUOffloadHooks() + base = torch.arange(20) + offset_view = torch.as_strided(base, (3, 4), (1, 3), 2) + + self.assertEqual(offset_view.storage_offset(), 2) + transfer_tensor, restore_view = hooks._transfer_view(offset_view) + restored = hooks.unpack((transfer_tensor.clone(), torch.device("cpu"), None, restore_view)) + + self.assertEqual(tuple(restored.shape), tuple(offset_view.shape)) + self.assertEqual(tuple(restored.stride()), tuple(offset_view.stride())) + self.assertTrue(torch.equal(restored, offset_view)) + + def test_cpu_offload_sparse_storage_views_keep_original_layout(self): + """Views with holes in storage are not flattened because storage-order transfer would lose layout.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import CPUOffloadHooks + + hooks = CPUOffloadHooks() + sparse_view = torch.arange(12).view(3, 4)[:, ::2] + + flat = hooks._flat_storage_view(sparse_view) + transfer_tensor, restore_view = hooks._transfer_view(sparse_view) + + self.assertIsNone(flat) + self.assertIs(transfer_tensor, sparse_view) + self.assertIsNone(restore_view) + + def test_cpu_offload_flat_transfer_key_ignores_dense_view_stride(self): + """Pinned bucket keys are based on transfer shape, not dense view logical stride.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import CPUOffloadHooks, _PinnedMemoryPool + + hooks = CPUOffloadHooks() + pool = _PinnedMemoryPool(max_buckets=2) + base = torch.empty(3, 4) + transposed = base.t() + + base_transfer, base_restore = hooks._transfer_view(base) + transposed_transfer, transposed_restore = hooks._transfer_view(transposed) + base_key = pool.key_for(base_transfer) + transposed_key = pool.key_for(transposed_transfer) + + self.assertEqual(base_key, transposed_key) + self.assertEqual(base_key.size, (12,)) + self.assertEqual(base_key.stride, (1,)) + self.assertEqual(base_restore.size, (3, 4)) + self.assertEqual(transposed_restore.size, (4, 3)) + self.assertEqual(transposed_restore.stride, (1, 4)) + + @unittest.skipIf(not torch.cuda.is_available(), "CUDA required for pinned offload test") + def test_activation_offload_pin_memory_bucket_limit(self): + """New shapes fall back to pageable CPU once the pinned bucket cap is reached.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + CPUOffloadHooks, + get_activation_offload_pin_memory_max_buckets, + get_activation_offload_pin_memory_stats, + reset_activation_offload_pin_memory_stats, + set_activation_offload_pin_memory_max_buckets, + ) + + original = get_activation_offload_pin_memory_max_buckets() + try: + set_activation_offload_pin_memory_max_buckets(0) + reset_activation_offload_pin_memory_stats() + set_activation_offload_pin_memory_max_buckets(1) + hooks = CPUOffloadHooks() + + first = hooks.pack(torch.randn(4, 4, device="cuda", requires_grad=True) * 2) + second = hooks.pack(torch.randn(8, 8, device="cuda", requires_grad=True) * 2) + + self.assertTrue(first[0].is_pinned()) + self.assertIsNotNone(first[2]) + self.assertIsNone(second[2]) + self.assertEqual(get_activation_offload_pin_memory_stats()["total_cap_misses"], 1) + finally: + set_activation_offload_pin_memory_max_buckets(original) + + @unittest.skipIf(not torch.cuda.is_available(), "CUDA required for copy stream offload test") + def test_cpu_offload_uses_copy_stream_ready_event(self): + """Pinned CUDA offload returns a ready event and restores dense view strides.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import CPUOffloadHooks + + hooks = CPUOffloadHooks() + tensor = (torch.arange(12, device="cuda", dtype=torch.float32, requires_grad=True) * 2).view(3, 4).t() + + packed = hooks.pack(tensor) + + self.assertGreaterEqual(len(packed), 5) + self.assertIsNotNone(packed[2]) + self.assertIsNotNone(packed[3]) + self.assertIsNotNone(packed[4]) + + unpacked = hooks.unpack(packed) + torch.cuda.synchronize() + + self.assertEqual(unpacked.device.type, "cuda") + self.assertEqual(tuple(unpacked.shape), tuple(tensor.shape)) + self.assertEqual(tuple(unpacked.stride()), tuple(tensor.stride())) + self.assertTrue(torch.equal(unpacked, tensor.detach())) + + @unittest.skipIf(not torch.cuda.is_available(), "CUDA required for copy stream pool test") + def test_activation_offload_copy_stream_pools_round_robin_by_direction(self): + """D2H offload and H2D prefetch use independent round-robin stream pools.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + _ACTIVATION_PREFETCH_RUNTIME, + _PINNED_MEMORY_POOL, + CPUOffloadHooks, + _OffloadedActivationRecord, + get_activation_offload_copy_stream_stats, + get_activation_offload_d2h_copy_stream_count, + get_activation_offload_h2d_prefetch_stream_count, + reset_activation_offload_copy_stream_stats, + reset_activation_offload_prefetch_stats, + set_activation_offload_d2h_copy_stream_count, + set_activation_offload_h2d_prefetch_stream_count, + ) + + original_d2h = get_activation_offload_d2h_copy_stream_count() + original_h2d = get_activation_offload_h2d_prefetch_stream_count() + try: + set_activation_offload_d2h_copy_stream_count(2) + set_activation_offload_h2d_prefetch_stream_count(2) + reset_activation_offload_prefetch_stats() + reset_activation_offload_copy_stream_stats() + + hooks = CPUOffloadHooks() + hooks.pack(torch.randn(32, 32, device="cuda", requires_grad=True) * 2) + hooks.pack(torch.randn(32, 32, device="cuda", requires_grad=True) * 3) + + pinned_a = torch.empty(32, 32, pin_memory=True) + pinned_b = torch.empty(32, 32, pin_memory=True) + pool_key = _PINNED_MEMORY_POOL.key_for(pinned_a) + for predictor_id, tensor in (("a", pinned_a), ("b", pinned_b)): + _ACTIVATION_PREFETCH_RUNTIME.register( + _OffloadedActivationRecord( + logical_id=predictor_id, + predictor_id=predictor_id, + generation=0, + tensor=tensor, + original_device=torch.device("cuda"), + pool_key=pool_key, + restore_view=None, + ready_event=None, + ) + ) + self.assertTrue(_ACTIVATION_PREFETCH_RUNTIME.prefetch(0, predictor_id)) + + torch.cuda.synchronize() + stats = get_activation_offload_copy_stream_stats() + d2h_devices = list(stats["d2h"]["devices"].values()) + h2d_devices = list(stats["h2d_prefetch"]["devices"].values()) + + self.assertEqual(d2h_devices[0]["uses"], [1, 1]) + self.assertEqual(h2d_devices[0]["uses"], [1, 1]) + finally: + reset_activation_offload_prefetch_stats() + set_activation_offload_d2h_copy_stream_count(original_d2h) + set_activation_offload_h2d_prefetch_stream_count(original_h2d) + + def test_activation_offload_prefetch_learns_stable_successors(self): + """Activation prefetch learns on stable predictor ids instead of unique payload ids.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + _ACTIVATION_PREFETCH_RUNTIME, + _OffloadedActivationRecord, + get_activation_offload_prefetch_enabled, + reset_activation_offload_prefetch_stats, + set_activation_offload_prefetch_enabled, + ) + + reset_activation_offload_prefetch_stats() + original_enabled = get_activation_offload_prefetch_enabled() + try: + set_activation_offload_prefetch_enabled(False) + for generation in range(2): + first = _OffloadedActivationRecord( + logical_id=f"{generation}:a:payload", + predictor_id="block.attn:0", + generation=generation, + tensor=torch.empty(1), + original_device=torch.device("cpu"), + pool_key=None, + restore_view=None, + ready_event=None, + ) + second = _OffloadedActivationRecord( + logical_id=f"{generation}:b:payload", + predictor_id="block.attn:1", + generation=generation, + tensor=torch.empty(1), + original_device=torch.device("cpu"), + pool_key=None, + restore_view=None, + ready_event=None, + ) + _ACTIVATION_PREFETCH_RUNTIME.register(first) + _ACTIVATION_PREFETCH_RUNTIME.register(second) + _ACTIVATION_PREFETCH_RUNTIME.consume(first) + _ACTIVATION_PREFETCH_RUNTIME.consume(second) + + stats = _ACTIVATION_PREFETCH_RUNTIME.snapshot() + self.assertEqual(_ACTIVATION_PREFETCH_RUNTIME.successors["block.attn:0"], "block.attn:1") + self.assertEqual(stats["learned_successors"], 1) + self.assertEqual(stats["transition_updates"], 1) + finally: + set_activation_offload_prefetch_enabled(original_enabled) + reset_activation_offload_prefetch_stats() + + def test_activation_offload_prefetch_retires_unconsumed_generation_records(self): + """Unpacked-only subsets should not strand pinned buffers in the prefetch runtime.""" + from simpletuner.helpers.training import offloaded_gradient_checkpointer as offload + + original_pool = offload._PINNED_MEMORY_POOL + original_runtime = offload._ACTIVATION_PREFETCH_RUNTIME + pool = offload._PinnedMemoryPool(max_buckets=1) + pool._allocate = lambda key: torch.empty_strided(key.size, key.stride, dtype=key.dtype, layout=key.layout) + runtime = offload._ActivationOffloadPrefetchRuntime() + try: + offload._PINNED_MEMORY_POOL = pool + offload._ACTIVATION_PREFETCH_RUNTIME = runtime + key = pool.key_for(torch.empty(2, 3)) + cpu_tensor = pool.checkout(key) + record = offload._OffloadedActivationRecord( + logical_id="0:unused:payload", + predictor_id="unused", + generation=0, + tensor=cpu_tensor, + original_device=torch.device("cuda" if torch.cuda.is_available() else "cpu"), + pool_key=key, + restore_view=None, + ready_event=None, + ) + + runtime.register(record) + runtime.saw_unpack_since_pack = True + + self.assertEqual(runtime.snapshot()["active_records"], 1) + self.assertEqual(runtime.next_generation_for_pack(), 1) + + snapshot = runtime.snapshot() + pool_snapshot = pool.snapshot() + self.assertEqual(snapshot["active_records"], 0) + self.assertTrue(record.cpu_released) + self.assertEqual(pool_snapshot["buckets"][0]["releases"], 1) + self.assertEqual(pool_snapshot["buckets"][0]["available_buffers"], 1) + finally: + offload._PINNED_MEMORY_POOL = original_pool + offload._ACTIVATION_PREFETCH_RUNTIME = original_runtime + + def test_activation_offload_prefetch_autotune_disables_worse_prefetch(self): + """Autotune disables prefetch when measured waits are not better than JIT restore.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + _ACTIVATION_PREFETCH_RUNTIME, + reset_activation_offload_prefetch_stats, + ) + + reset_activation_offload_prefetch_stats() + try: + for _ in range(8): + _ACTIVATION_PREFETCH_RUNTIME.record_jit_restore_ms(1.0) + _ACTIVATION_PREFETCH_RUNTIME.record_prefetch_wait_ms(1.1) + + stats = _ACTIVATION_PREFETCH_RUNTIME.snapshot() + self.assertTrue(stats["autotune_disabled"]) + self.assertEqual(stats["autotune_decision"], "jit") + finally: + reset_activation_offload_prefetch_stats() + class TestGradientCheckpointingBackend(unittest.TestCase): """Tests for the gradient checkpointing backend module.""" + def test_trainer_prefetch_autotune_uses_model_predict_and_discards_gradients(self): + """Trainer-level prefetch autotune probes the real prediction/loss path without retaining grads.""" + from simpletuner.helpers.training.offloaded_gradient_checkpointer import ( + activation_offload_prefetch_autotune_decision, + reset_activation_offload_prefetch_stats, + ) + from simpletuner.helpers.training.trainer import Trainer + + class FakeAccelerator: + def __init__(self): + self.backward_calls = 0 + + def backward(self, loss): + self.backward_calls += 1 + loss.backward() + + class FakeModel: + def __init__(self, component): + self.component = component + + def get_trained_component(self, unwrap_model=False): + return self.component + + def loss_with_logs(self, prepared_batch, model_output, apply_conditioning_mask=True): + return model_output.square().mean(), {} + + def auxiliary_loss(self, prepared_batch, model_output, loss): + return loss, {} + + trainer = Trainer.__new__(Trainer) + trainer.config = SimpleNamespace( + gradient_checkpointing_offload_prefetch=True, + gradient_checkpointing_offload_attention=True, + disable_accelerator=False, + distillation_method=None, + ) + trainer.probe_component = nn.Linear(2, 2) + trainer.model = FakeModel(trainer.probe_component) + trainer.optimizer = torch.optim.SGD(trainer.probe_component.parameters(), lr=0.1) + trainer.sidecar_optimizer = None + trainer.accelerator = FakeAccelerator() + trainer.predict_calls = 0 + + def model_predict(prepared_batch): + trainer.predict_calls += 1 + return trainer.probe_component(prepared_batch["x"]) + + trainer.model_predict = model_predict + prepared_batch = {"x": torch.ones(1, 2)} + + reset_activation_offload_prefetch_stats() + with ( + mock.patch("simpletuner.helpers.training.trainer.torch.cuda.is_available", return_value=True), + mock.patch("simpletuner.helpers.training.trainer.torch.cuda.synchronize"), + mock.patch("simpletuner.helpers.training.trainer.torch.cuda.get_rng_state_all", return_value=[]), + mock.patch("simpletuner.helpers.training.trainer.torch.cuda.set_rng_state_all"), + ): + trainer._maybe_autotune_activation_offload_prefetch(prepared_batch) + + self.assertEqual(trainer.predict_calls, 3) + self.assertEqual(trainer.accelerator.backward_calls, 3) + self.assertIsNone(trainer.probe_component.weight.grad) + self.assertIsNone(trainer.probe_component.bias.grad) + self.assertIn(activation_offload_prefetch_autotune_decision(), {"jit", "prefetch"}) + reset_activation_offload_prefetch_stats() + def test_set_checkpoint_backend(self): """Test that checkpoint backend can be set.""" from simpletuner.helpers.training.gradient_checkpointing_interval import ( @@ -263,6 +759,123 @@ def test_checkpoint_function_produces_correct_output(self): finally: set_checkpoint_backend(original_backend) + def test_checkpoint_sequential_state_matches_direct_gradients(self): + """Test segmented checkpointing over a tuple-carrying block sequence.""" + from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state + + class TupleBlock(nn.Module): + def __init__(self, dim: int): + super().__init__() + self.x_proj = nn.Linear(dim, dim) + self.y_proj = nn.Linear(dim, dim) + + def forward(self, x, y): + next_x = torch.relu(self.x_proj(x) + y) + next_y = torch.relu(self.y_proj(y) + next_x) + return next_x, next_y + + direct_blocks = nn.ModuleList([TupleBlock(8) for _ in range(4)]) + checkpointed_blocks = nn.ModuleList([TupleBlock(8) for _ in range(4)]) + checkpointed_blocks.load_state_dict(direct_blocks.state_dict()) + + direct_x = torch.randn(2, 8, requires_grad=True) + direct_y = torch.randn(2, 8, requires_grad=True) + checkpointed_x = direct_x.detach().clone().requires_grad_(True) + checkpointed_y = direct_y.detach().clone().requires_grad_(True) + + x, y = direct_x, direct_y + for block in direct_blocks: + x, y = block(x, y) + direct_loss = x.sum() + y.sum() + direct_loss.backward() + + def run_block(_index, block, x, y): + return block(x, y) + + x, y = checkpoint_sequential_state( + list(checkpointed_blocks), + 2, + (checkpointed_x, checkpointed_y), + run_block, + torch.utils.checkpoint.checkpoint, + {"use_reentrant": False}, + ) + checkpointed_loss = x.sum() + y.sum() + checkpointed_loss.backward() + + self.assertTrue(torch.allclose(direct_x.grad, checkpointed_x.grad, atol=1e-6)) + self.assertTrue(torch.allclose(direct_y.grad, checkpointed_y.grad, atol=1e-6)) + for direct_block, checkpointed_block in zip(direct_blocks, checkpointed_blocks): + for direct_param, checkpointed_param in zip(direct_block.parameters(), checkpointed_block.parameters()): + self.assertTrue(torch.allclose(direct_param.grad, checkpointed_param.grad, atol=1e-6)) + + def test_checkpoint_sequential_state_uses_contiguous_chunks(self): + """Test that segment_size controls contiguous chunk boundaries.""" + from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state + + calls = [] + + def checkpoint_fn(function, *args, **_kwargs): + calls.append("checkpoint") + return function(*args) + + def run_block(index, block, x): + calls.append(index) + return x + block + + (result,) = checkpoint_sequential_state( + [1, 2, 3, 4, 5], + 2, + (torch.tensor(0),), + run_block, + checkpoint_fn, + {"use_reentrant": False}, + ) + + self.assertEqual(result.item(), 15) + self.assertEqual(calls, ["checkpoint", 0, 1, "checkpoint", 2, 3, "checkpoint", 4]) + + def test_checkpoint_sequential_state_segment_stride_runs_gaps_without_checkpoint(self): + """Test that segment_stride leaves deterministic eager gaps between chunks.""" + from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state + + calls = [] + + def checkpoint_fn(function, *args, **_kwargs): + calls.append("checkpoint") + return function(*args) + + def run_block(index, block, x): + calls.append(index) + return x + block + + (result,) = checkpoint_sequential_state( + [1, 2, 3, 4, 5, 6], + 2, + (torch.tensor(0),), + run_block, + checkpoint_fn, + {"use_reentrant": False}, + segment_stride=4, + ) + + self.assertEqual(result.item(), 21) + self.assertEqual(calls, ["checkpoint", 0, 1, 2, 3, "checkpoint", 4, 5]) + + def test_checkpoint_sequential_state_rejects_overlapping_stride(self): + """Test that overlapping segment schedules are rejected.""" + from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state + + with self.assertRaisesRegex(ValueError, "segment_stride"): + checkpoint_sequential_state( + [1, 2], + 2, + (torch.tensor(0),), + lambda _index, block, x: x + block, + lambda function, *args, **_kwargs: function(*args), + segment_stride=1, + ) + class TestConfigFieldIntegration(unittest.TestCase): """Tests for the configuration field integration.""" @@ -281,6 +894,42 @@ def test_gradient_checkpointing_backend_field_exists(self): self.assertIn({"value": "unsloth", "label": "Unsloth layer (CPU offload)"}, field.choices) self.assertIn({"value": "unsloth-ffn", "label": "Unsloth FFN-only (CPU offload)"}, field.choices) + def test_gradient_checkpointing_offload_attention_field_exists(self): + """Test that the attention activation offload field is registered.""" + from simpletuner.simpletuner_sdk.server.services.field_registry import FieldRegistry + + registry = FieldRegistry() + field = registry.get_field("gradient_checkpointing_offload_attention") + + self.assertIsNotNone(field) + self.assertEqual(field.default_value, False) + self.assertEqual(field.arg_name, "--gradient_checkpointing_offload_attention") + self.assertEqual(field.dependencies, []) + + def test_gradient_checkpointing_offload_pin_memory_max_buckets_field_exists(self): + """Test that the attention offload pinned bucket field is registered.""" + from simpletuner.simpletuner_sdk.server.services.field_registry import FieldRegistry + + registry = FieldRegistry() + field = registry.get_field("gradient_checkpointing_offload_pin_memory_max_buckets") + + self.assertIsNotNone(field) + self.assertEqual(field.default_value, 12) + self.assertEqual(field.arg_name, "--gradient_checkpointing_offload_pin_memory_max_buckets") + self.assertEqual(len(field.dependencies), 1) + self.assertEqual(field.dependencies[0].field, "gradient_checkpointing_offload_attention") + + def test_gradient_checkpointing_offload_prefetch_field_exists(self): + """Test that the attention offload prefetch field is registered.""" + from simpletuner.simpletuner_sdk.server.services.field_registry import FieldRegistry + + registry = FieldRegistry() + field = registry.get_field("gradient_checkpointing_offload_prefetch") + + self.assertIsNotNone(field) + self.assertEqual(field.default_value, False) + self.assertEqual(field.arg_name, "--gradient_checkpointing_offload_prefetch") + def test_gradient_checkpointing_backend_validation(self): """Test that invalid backend values are rejected.""" from simpletuner.simpletuner_sdk.server.services.field_registry import FieldRegistry @@ -301,86 +950,16 @@ def test_gradient_checkpointing_backend_validation(self): self.assertIn("unsloth", choices_rule.value) self.assertIn("unsloth-ffn", choices_rule.value) + def test_gradient_checkpointing_segment_stride_field_exists(self): + """Test that the segmented checkpointing stride field is registered.""" + from simpletuner.simpletuner_sdk.server.services.field_registry import FieldRegistry -class TestTransformerBackendAttribute(unittest.TestCase): - """Tests that transformer models have the backend attribute and setter.""" - - def test_flux_transformer_has_backend_attribute(self): - """Test that FluxTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.flux.transformer import FluxTransformer2DModel - - self.assertTrue(hasattr(FluxTransformer2DModel, "set_gradient_checkpointing_backend")) - self.assertTrue(getattr(FluxTransformer2DModel, "_supports_ffn_gradient_checkpointing", False)) - - def test_flux_blocks_support_ffn_checkpoint_scope(self): - """Test that Flux blocks preserve output values with FFN-only checkpointing.""" - from simpletuner.helpers.models.flux.transformer import FluxSingleTransformerBlock, FluxTransformerBlock - - double_block = FluxTransformerBlock(dim=16, num_attention_heads=2, attention_head_dim=8).train() - hidden = torch.randn(2, 4, 16, requires_grad=True) - encoder_hidden = torch.randn(2, 3, 16, requires_grad=True) - temb = torch.randn(2, 16) - - expected_encoder, expected_hidden = double_block(hidden, encoder_hidden, temb) - actual_encoder, actual_hidden = double_block( - hidden, - encoder_hidden, - temb, - checkpoint_ffn=True, - checkpoint_fn=torch.utils.checkpoint.checkpoint, - ) - self.assertTrue(torch.allclose(expected_encoder, actual_encoder, atol=1e-6)) - self.assertTrue(torch.allclose(expected_hidden, actual_hidden, atol=1e-6)) - - single_block = FluxSingleTransformerBlock(dim=16, num_attention_heads=2, attention_head_dim=8).train() - hidden = torch.randn(2, 7, 16, requires_grad=True) - temb = torch.randn(2, 16) - - expected_hidden = single_block(hidden, temb) - actual_hidden = single_block( - hidden, - temb, - checkpoint_ffn=True, - checkpoint_fn=torch.utils.checkpoint.checkpoint, - ) - self.assertTrue(torch.allclose(expected_hidden, actual_hidden, atol=1e-6)) - - def test_sana_transformer_has_backend_attribute(self): - """Test that SanaTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.sana.transformer import SanaTransformer2DModel - - self.assertTrue(hasattr(SanaTransformer2DModel, "set_gradient_checkpointing_backend")) - - def test_sd3_transformer_has_backend_attribute(self): - """Test that SD3Transformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.sd3.transformer import SD3Transformer2DModel - - self.assertTrue(hasattr(SD3Transformer2DModel, "set_gradient_checkpointing_backend")) - - def test_chroma_transformer_has_backend_attribute(self): - """Test that ChromaTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.chroma.transformer import ChromaTransformer2DModel - - self.assertTrue(hasattr(ChromaTransformer2DModel, "set_gradient_checkpointing_backend")) - - def test_auraflow_transformer_has_backend_attribute(self): - """Test that AuraFlowTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.auraflow.transformer import AuraFlowTransformer2DModel - - self.assertTrue(hasattr(AuraFlowTransformer2DModel, "set_gradient_checkpointing_backend")) - - def test_mageflow_transformer_has_backend_attribute(self): - """Test that MageFlowTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.mageflow.transformer import MageFlowTransformer2DModel - - self.assertTrue(hasattr(MageFlowTransformer2DModel, "set_gradient_checkpointing_backend")) - self.assertTrue(getattr(MageFlowTransformer2DModel, "_supports_ffn_gradient_checkpointing", False)) - - def test_qwen_image_transformer_has_backend_attribute(self): - """Test that QwenImageTransformer2DModel has gradient_checkpointing_backend.""" - from simpletuner.helpers.models.qwen_image.transformer import QwenImageTransformer2DModel + registry = FieldRegistry() + field = registry.get_field("gradient_checkpointing_segment_stride") - self.assertTrue(hasattr(QwenImageTransformer2DModel, "set_gradient_checkpointing_backend")) + self.assertIsNotNone(field) + self.assertEqual(field.default_value, None) + self.assertEqual(field.arg_name, "--gradient_checkpointing_segment_stride") if __name__ == "__main__": diff --git a/tests/test_hidream_model.py b/tests/test_hidream_model.py index e1e515e82..781755278 100644 --- a/tests/test_hidream_model.py +++ b/tests/test_hidream_model.py @@ -67,6 +67,57 @@ def _forward(**kwargs): transformer_kwargs = self.model.model.call_args.kwargs self.assertTrue(torch.equal(transformer_kwargs["timesteps"], prepared_batch["timesteps"])) + def test_check_user_config_keeps_bundled_quanto_text_encoder_with_quanto_base(self): + self.model.config = SimpleNamespace( + base_model_precision="int8-quanto", + text_encoder_4_precision="int4-quanto", + tokenizer_max_length=0, + i_know_what_i_am_doing=False, + aspect_bucket_alignment=32, + ) + + self.model.check_user_config() + + self.assertEqual(self.model.config.text_encoder_4_precision, "int4-quanto") + + def test_check_user_config_disables_bundled_quanto_text_encoder_for_torchao_base(self): + self.model.config = SimpleNamespace( + base_model_precision="fp8-torchao", + text_encoder_4_precision="int4-quanto", + tokenizer_max_length=0, + i_know_what_i_am_doing=False, + aspect_bucket_alignment=32, + ) + + self.model.check_user_config() + + self.assertEqual(self.model.config.text_encoder_4_precision, "no_change") + + def test_check_user_config_disables_bundled_quanto_text_encoder_for_sdnq_base(self): + self.model.config = SimpleNamespace( + base_model_precision="int8-sdnq", + text_encoder_4_precision="int4-quanto", + tokenizer_max_length=0, + i_know_what_i_am_doing=False, + aspect_bucket_alignment=32, + ) + + self.model.check_user_config() + + self.assertEqual(self.model.config.text_encoder_4_precision, "no_change") + + def test_check_user_config_rejects_non_default_mixed_text_encoder_backend(self): + self.model.config = SimpleNamespace( + base_model_precision="fp8-torchao", + text_encoder_4_precision="int8-quanto", + tokenizer_max_length=0, + i_know_what_i_am_doing=False, + aspect_bucket_alignment=32, + ) + + with self.assertRaisesRegex(ValueError, "cannot mix base model precision"): + self.model.check_user_config() + if __name__ == "__main__": unittest.main() diff --git a/tests/test_hunyuanvideo_model.py b/tests/test_hunyuanvideo_model.py index c9d356bca..e0c01565f 100644 --- a/tests/test_hunyuanvideo_model.py +++ b/tests/test_hunyuanvideo_model.py @@ -76,10 +76,82 @@ def test_load_text_encoder_registers_both_hunyuan_encoders_for_device_management self.assertIs(model.text_encoder_2, byt5_model) self.assertIs(model.get_text_encoder(1), byt5_model) + def test_load_text_encoder_prefers_qwen_text_encoder_override(self): + model = HunyuanVideo.__new__(HunyuanVideo) + model.accelerator = SimpleNamespace(device=torch.device("cuda:0")) + model.config = SimpleNamespace( + qwen_text_encoder_model_name_or_path="custom/qwen", + hunyuan_text_encoder_path="legacy/qwen", + glyph_byt5_repo="glyph/repo", + glyph_byt5_fallback_repo="glyph/fallback", + ) + model._ramtorch_text_encoders_requested = MagicMock(return_value=False) + model._ramtorch_text_encoder_percent = MagicMock(return_value=1.0) + model._apply_ramtorch_layers = MagicMock() + + qwen_tokenizer = MagicMock() + byt5_tokenizer = MagicMock() + text_encoder = MagicMock() + text_encoder.to.return_value = text_encoder + byt5_model = MagicMock() + byt5_model.to.return_value = byt5_model + + with ( + patch( + "simpletuner.helpers.models.hunyuanvideo.model.Qwen2Tokenizer.from_pretrained", + return_value=qwen_tokenizer, + ) as mock_qwen_tokenizer, + patch( + "simpletuner.helpers.models.hunyuanvideo.model.Qwen2_5_VLTextModel.from_pretrained", + return_value=text_encoder, + ) as mock_qwen_text_encoder, + patch( + "simpletuner.helpers.models.hunyuanvideo.model.ByT5Tokenizer.from_pretrained", + return_value=byt5_tokenizer, + ), + patch( + "simpletuner.helpers.models.hunyuanvideo.model.T5EncoderModel.from_pretrained", + return_value=byt5_model, + ), + patch( + "simpletuner.helpers.models.hunyuanvideo.model.hf_hub_download", + side_effect=RuntimeError("no glyph checkpoint"), + ), + ): + model.load_text_encoder(move_to_device=True) + + mock_qwen_tokenizer.assert_called_once_with("custom/qwen") + mock_qwen_text_encoder.assert_called_once_with("custom/qwen", torch_dtype=torch.bfloat16) + def test_model_supports_crepa_self_flow(self): model = HunyuanVideo.__new__(HunyuanVideo) self.assertTrue(model.supports_crepa_self_flow()) + def test_check_user_config_maps_upstream_repo_to_diffusers_flavour_repo(self): + model = HunyuanVideo.__new__(HunyuanVideo) + model.config = SimpleNamespace( + model_flavour="t2v-480p", + pretrained_model_name_or_path=HunyuanVideo.UPSTREAM_MODEL_REPO, + ) + + model.check_user_config() + + self.assertEqual( + model.config.pretrained_model_name_or_path, + HunyuanVideo.HUGGINGFACE_PATHS["t2v-480p"], + ) + + def test_check_user_config_preserves_explicit_custom_model_path(self): + model = HunyuanVideo.__new__(HunyuanVideo) + model.config = SimpleNamespace( + model_flavour="t2v-480p", + pretrained_model_name_or_path="local-or-hub/custom-hunyuan-diffusers", + ) + + model.check_user_config() + + self.assertEqual(model.config.pretrained_model_name_or_path, "local-or-hub/custom-hunyuan-diffusers") + def test_prepare_crepa_self_flow_batch_builds_tokenwise_student_and_teacher_views(self): model = HunyuanVideo.__new__(HunyuanVideo) model.accelerator = SimpleNamespace(device=torch.device("cpu")) diff --git a/tests/test_longcat_image_model.py b/tests/test_longcat_image_model.py index 090d19971..b6df0756b 100644 --- a/tests/test_longcat_image_model.py +++ b/tests/test_longcat_image_model.py @@ -5,6 +5,7 @@ import torch from simpletuner.helpers.models.longcat_image.model import LongCatImage +from simpletuner.helpers.models.longcat_image.pipeline import LongCatImagePipeline class LongCatImageModelTests(unittest.TestCase): @@ -18,6 +19,18 @@ def setUp(self): def test_model_supports_crepa_self_flow(self): self.assertTrue(self.model.supports_crepa_self_flow()) + def test_pipeline_allows_missing_image_encoder(self): + pipeline = LongCatImagePipeline( + scheduler=MagicMock(), + vae=MagicMock(), + text_encoder=MagicMock(), + tokenizer=MagicMock(), + text_processor=MagicMock(), + transformer=MagicMock(), + ) + + self.assertIsNone(pipeline.image_encoder) + def test_prepare_crepa_self_flow_batch_creates_tokenwise_timesteps(self): self.model.accelerator = SimpleNamespace(device=torch.device("cpu")) self.model.config = SimpleNamespace( diff --git a/tests/test_lora_format.py b/tests/test_lora_format.py new file mode 100644 index 000000000..f76f81b6b --- /dev/null +++ b/tests/test_lora_format.py @@ -0,0 +1,98 @@ +import unittest + +import torch + +from simpletuner.helpers.training.lora_format import ( + PEFTLoRAFormat, + convert_diffusers_to_comfyui, + convert_diffusers_to_comfyui_sd_lora, + detect_state_dict_format, +) + +DOWN = torch.full((4, 8), 1.0) +UP = torch.full((8, 4), 2.0) + + +def _diffusers_named(module_key): + return { + f"{module_key}.lora.down.weight": DOWN, + f"{module_key}.lora.up.weight": UP, + } + + +def _peft_named(module_key): + return { + f"{module_key}.lora_A.weight": DOWN, + f"{module_key}.lora_B.weight": UP, + } + + +SPELLINGS = (("diffusers", _diffusers_named), ("peft", _peft_named)) + + +class ConvertDiffusersToComfyUITests(unittest.TestCase): + MODULE = "transformer.blocks.0.attn.to_q" + CONVERTED = "diffusion_model.blocks.0.attn.to_q" + ALPHA_KEY = "diffusion_model.blocks.0.attn.to_q.alpha" + METADATA = {"lora_alpha": 16} + + def _convert(self, state_dict): + return convert_diffusers_to_comfyui(state_dict, adapter_metadata=self.METADATA) + + def test_both_spellings_convert_to_the_same_keys(self): + converted = {name: set(self._convert(build(self.MODULE))) for name, build in SPELLINGS} + self.assertEqual(converted["diffusers"], converted["peft"]) + + def test_alpha_is_emitted_for_both_spellings(self): + for name, build in SPELLINGS: + with self.subTest(spelling=name): + converted = self._convert(build(self.MODULE)) + self.assertIn(self.ALPHA_KEY, converted) + self.assertEqual(float(converted[self.ALPHA_KEY]), float(self.METADATA["lora_alpha"])) + + def test_down_and_up_weights_keep_their_roles(self): + for name, build in SPELLINGS: + with self.subTest(spelling=name): + converted = self._convert(build(self.MODULE)) + self.assertTrue(torch.equal(converted[f"{self.CONVERTED}.lora_A.weight"], DOWN)) + self.assertTrue(torch.equal(converted[f"{self.CONVERTED}.lora_B.weight"], UP)) + + +class ConvertDiffusersToComfyUISDLoraTests(unittest.TestCase): + MODULE = "unet.down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k" + ALPHA_KEY = "lora_unet_down_blocks_1_attentions_0_transformer_blocks_0_attn1_to_k.alpha" + METADATA = {"lora_alpha": 8} + + def _convert(self, state_dict): + return convert_diffusers_to_comfyui_sd_lora(state_dict, adapter_metadata=self.METADATA) + + def test_both_spellings_convert_to_the_same_keys(self): + converted = {name: set(self._convert(build(self.MODULE))) for name, build in SPELLINGS} + self.assertEqual(converted["diffusers"], converted["peft"]) + + def test_alpha_is_emitted_for_both_spellings(self): + for name, build in SPELLINGS: + with self.subTest(spelling=name): + converted = self._convert(build(self.MODULE)) + self.assertIn(self.ALPHA_KEY, converted) + self.assertEqual(float(converted[self.ALPHA_KEY]), float(self.METADATA["lora_alpha"])) + + +class DetectStateDictFormatTests(unittest.TestCase): + MODULE = "transformer.blocks.0.attn.to_q" + + def test_peft_named_dict_is_reported_as_diffusers(self): + self.assertEqual(detect_state_dict_format(_peft_named(self.MODULE)), PEFTLoRAFormat.DIFFUSERS) + + def test_diffusers_named_dict_is_reported_as_diffusers(self): + self.assertEqual(detect_state_dict_format(_diffusers_named(self.MODULE)), PEFTLoRAFormat.DIFFUSERS) + + def test_diffusion_model_peft_named_dict_is_reported_as_comfyui(self): + self.assertEqual( + detect_state_dict_format(_peft_named("diffusion_model.blocks.0.attn.to_q")), + PEFTLoRAFormat.COMFYUI, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_lora_metadata.py b/tests/test_lora_metadata.py index aebd43551..eb02cbf1f 100644 --- a/tests/test_lora_metadata.py +++ b/tests/test_lora_metadata.py @@ -41,7 +41,8 @@ def __init__(self, name): class _DummyTrainedComponent(_DummyBaseModule): - pass + def parameters(self): + return [] class _DummyControlNet(_DummyTrainedComponent): @@ -131,6 +132,14 @@ def save_lora_weights( save_function(packed, filename) +class _DummyDiffusersSavePipeline(_DummySavePipeline): + @classmethod + def save_lora_weights(cls, save_directory, transformer_lora_layers=None, save_function=None, **kwargs): + filename = os.path.join(save_directory, kwargs.get("weight_name") or "pytorch_lora_weights.safetensors") + packed = {f"transformer.{key}": value for key, value in (transformer_lora_layers or {}).items()} + (save_function or save_file)(packed, filename) + + class _DummySDXLSavePipeline: @classmethod def save_lora_weights( @@ -165,6 +174,20 @@ class _DummySDXLSaveModel(_DummySaveModel): PIPELINE_CLASSES = {PipelineTypes.TEXT2IMG: _DummySDXLSavePipeline, PipelineTypes.CONTROLNET: _DummySDXLSavePipeline} +class _DummyArtifactModel(_DummyModel): + """Routes _save_lora through the real ModelFoundation.save_lora_weights so a file is written.""" + + PIPELINE_CLASSES = { + PipelineTypes.TEXT2IMG: _DummyDiffusersSavePipeline, + PipelineTypes.CONTROLNET: _DummyDiffusersSavePipeline, + } + + def __init__(self, trained_component, text_encoder=None): + super().__init__(trained_component=trained_component, text_encoder=text_encoder) + self.config = SimpleNamespace(model_family="sdxl", lora_format="diffusers", controlnet=False) + self.save_lora_weights = lambda *args, **kwargs: ModelFoundation.save_lora_weights(self, *args, **kwargs) + + _ema_stub = SimpleNamespace( store=lambda *args, **kwargs: None, copy_to=lambda *args, **kwargs: None, restore=lambda *args, **kwargs: None ) @@ -196,6 +219,27 @@ def _make_manager(self, args_overrides=None, text_encoder=None): ) return manager, model, trained_component + def test_manager_rejects_invalid_default_pipeline_type(self): + args = SimpleNamespace( + use_ema=False, + model_type="lora", + lora_type="standard", + controlnet=False, + validation_using_datasets=False, + ) + trained_component = _DummyTrainedComponent("transformer") + model = _DummyModel(trained_component=trained_component) + model.DEFAULT_PIPELINE_TYPE = None + + with self.assertRaisesRegex(ValueError, "DEFAULT_PIPELINE_TYPE"): + SaveHookManager( + args=args, + model=model, + ema_model=_ema_stub, + accelerator=_DummyAccelerator(), + use_deepspeed_optimizer=False, + ) + def test_materialize_state_dict_for_save_expands_dtensor_like_values(self): class DTensor: def __init__(self): @@ -414,6 +458,189 @@ def test_save_hook_handles_controlnet_metadata(self): self.assertNotIn("transformer_lora_adapter_metadata", kwargs) self.assertIn("text_encoder_lora_adapter_metadata", kwargs) + def test_save_hook_writes_denoiser_lora_in_peft_naming(self): + manager, model, trained_component = self._make_manager() + + lora_state = { + "transformer_blocks.0.attn.to_q.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.attn.to_q.lora_B.weight": torch.zeros(8, 4), + "transformer_blocks.0.ff.up.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.ff.up.lora_B.weight": torch.zeros(8, 4), + } + with patch("simpletuner.helpers.training.save_hooks.get_peft_model_state_dict", return_value=lora_state): + with tempfile.TemporaryDirectory() as tmpdir: + manager._save_lora(models=[trained_component], weights=[object()], output_dir=tmpdir) + + written = model.save_lora_weights.call_args.kwargs["transformer_lora_layers"] + self.assertEqual(set(written), set(lora_state)) + for key in written: + self.assertTrue(key.endswith((".lora_A.weight", ".lora_B.weight")), key) + self.assertNotIn(".lora.down", key) + self.assertNotIn(".lora.up", key) + + def test_save_hook_never_writes_a_mixed_naming_scheme(self): + lora_state = { + "transformer_blocks.0.attn.to_q.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.attn.to_q.lora_B.weight": torch.zeros(8, 4), + "transformer_blocks.0.ff.up.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.ff.up.lora_B.weight": torch.zeros(8, 4), + } + + for use_ema in (False, True): + with self.subTest(use_ema=use_ema): + manager, model, trained_component = self._make_manager(args_overrides={"use_ema": use_ema}) + with patch("simpletuner.helpers.training.save_hooks.get_peft_model_state_dict", return_value=lora_state): + with tempfile.TemporaryDirectory() as tmpdir: + manager._save_lora(models=[trained_component], weights=[object()], output_dir=tmpdir) + + calls = model.save_lora_weights.call_args_list + self.assertEqual(len(calls), 2 if use_ema else 1) + for call in calls: + written = call.kwargs["transformer_lora_layers"] + schemes = {"peft" if (".lora_A." in key or ".lora_B." in key) else "diffusers" for key in written} + self.assertEqual(schemes, {"peft"}, written) + + def test_saved_denoiser_file_uses_peft_naming(self): + args = SimpleNamespace( + use_ema=False, + model_type="lora", + lora_type="standard", + controlnet=False, + validation_using_datasets=False, + model_family="sdxl", + model_flavour="base-1.0", + tracker_run_name="run-name", + resolution=1024, + ) + trained_component = _DummyTrainedComponent("transformer") + manager = SaveHookManager( + args=args, + model=_DummyArtifactModel(trained_component=trained_component), + ema_model=_ema_stub, + accelerator=_DummyAccelerator(), + use_deepspeed_optimizer=False, + ) + + lora_state = { + "transformer_blocks.0.attn.to_q.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.attn.to_q.lora_B.weight": torch.zeros(8, 4), + "transformer_blocks.0.ff.up.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.ff.up.lora_B.weight": torch.zeros(8, 4), + } + with patch("simpletuner.helpers.training.save_hooks.get_peft_model_state_dict", return_value=lora_state): + with tempfile.TemporaryDirectory() as tmpdir: + manager._save_lora(models=[trained_component], weights=[object()], output_dir=tmpdir) + lora_path = os.path.join(tmpdir, "pytorch_lora_weights.safetensors") + with safe_open(lora_path, framework="pt", device="cpu") as handle: + written = list(handle.keys()) + + self.assertEqual(len(written), len(lora_state)) + for key in written: + self.assertTrue(key.endswith((".lora_A.weight", ".lora_B.weight")), key) + + def test_save_hook_peft_output_loads_through_diffusers_with_model_prefix(self): + from diffusers.loaders.peft import PeftAdapterMixin + from peft import LoraConfig, inject_adapter_in_model + from peft.utils import get_peft_model_state_dict + + class Block(torch.nn.Module): + def __init__(self): + super().__init__() + self.to_gate = torch.nn.Linear(8, 8, bias=False) + self.to_q = torch.nn.Linear(8, 8, bias=False) + + class Denoiser(PeftAdapterMixin, torch.nn.Module): + def __init__(self): + super().__init__() + self.model = torch.nn.Module() + self.model.transformer = torch.nn.Module() + self.model.transformer.foo = Block() + + class Pipeline: + pass + + class Model: + MODEL_CLASS = Denoiser + MODEL_SUBFOLDER = "model" + PIPELINE_CLASSES = {PipelineTypes.TEXT2IMG: Pipeline} + + def __init__(self, trained_component): + self.trained_component = trained_component + self.saved_lora_layers = None + + def get_trained_component(self, unwrap_model=False, base_model=False): + return self.trained_component + + def get_text_encoder(self, index): + return None + + def save_lora_weights(self, output_dir, **kwargs): + self.saved_lora_layers = kwargs["model_lora_layers"] + + def make_trained_component(): + component = Denoiser() + inject_adapter_in_model( + LoraConfig(r=2, lora_alpha=2, target_modules=["to_gate", "to_q"]), + component, + adapter_name="default", + ) + return component + + def load_with_diffusers_from_safetensors(state_dict): + target = Denoiser() + with tempfile.TemporaryDirectory() as tmpdir: + save_file(state_dict, os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + target.load_lora_adapter( + tmpdir, + prefix=None, + weight_name="pytorch_lora_weights.safetensors", + use_safetensors=True, + adapter_name="loaded", + ) + return sorted(name for name, module in target.named_modules() if hasattr(module, "lora_A")) + + expected_modules = ["model.transformer.foo.to_gate", "model.transformer.foo.to_q"] + raw_peft_state = get_peft_model_state_dict(make_trained_component()) + self.assertEqual(load_with_diffusers_from_safetensors(raw_peft_state), expected_modules) + + trained_component = make_trained_component() + model = Model(trained_component) + manager = SaveHookManager( + args=SimpleNamespace(use_ema=False, controlnet=False, validation_using_datasets=False), + model=model, + ema_model=_ema_stub, + accelerator=_DummyAccelerator(), + use_deepspeed_optimizer=False, + ) + + with tempfile.TemporaryDirectory() as tmpdir: + manager._save_lora(models=[trained_component], weights=[object()], output_dir=tmpdir, write=True) + + self.assertEqual(load_with_diffusers_from_safetensors(model.saved_lora_layers), expected_modules) + + def test_legacy_mixed_checkpoint_normalises_for_resume(self): + from diffusers.utils import convert_unet_state_dict_to_peft + + mixed = { + "transformer_blocks.0.attn.to_gate.lora_A.weight": torch.zeros(4, 8), + "transformer_blocks.0.attn.to_gate.lora_B.weight": torch.zeros(8, 4), + "transformer_blocks.0.attn.to_q.lora.down.weight": torch.zeros(4, 8), + "transformer_blocks.0.attn.to_q.lora.up.weight": torch.zeros(8, 4), + } + + # safetensors sorts keys, and diffusers' loader gates conversion on the first one alone. + self.assertIn("lora_A", sorted(mixed)[0]) + + converted = convert_unet_state_dict_to_peft(mixed) + + self.assertEqual(len(converted), len(mixed)) + for key in converted: + self.assertTrue(key.endswith((".lora_A.weight", ".lora_B.weight")), key) + self.assertEqual( + {key.rsplit(".lora_", 1)[0] for key in converted}, + {"transformer_blocks.0.attn.to_gate", "transformer_blocks.0.attn.to_q"}, + ) + def test_models_spec_metadata_written_to_lora_file(self): manager, _, _ = self._make_manager( args_overrides={ diff --git a/tests/test_mageflow_model.py b/tests/test_mageflow_model.py index b2f9ae73a..7bf3e4119 100644 --- a/tests/test_mageflow_model.py +++ b/tests/test_mageflow_model.py @@ -1,5 +1,6 @@ import inspect import json +import tempfile import unittest from pathlib import Path from types import SimpleNamespace @@ -8,8 +9,8 @@ import torch from simpletuner.helpers.acceleration import AccelerationBackend -from simpletuner.helpers.models.common import PipelineTypes -from simpletuner.helpers.models.mageflow.pipeline import MageFlowPipeline, _mageflow_velocity +from simpletuner.helpers.models.common import ImageModelFoundation, PipelineTypes +from simpletuner.helpers.models.mageflow.pipeline import MageFlowPipeline, _load_scheduler, _mageflow_velocity from simpletuner.helpers.models.mageflow.pipeline_edit import MageFlowEditPipeline from simpletuner.helpers.models.mageflow.transformer import MageFlowTransformer2DModel from simpletuner.helpers.models.mageflow.vendor.models.modules import _attn_backend as mageflow_attn_backend @@ -56,6 +57,7 @@ def test_registry_resolves_mageflow(self): self.assertIn("edit-turbo", model_cls.get_flavour_choices()) self.assertFalse(model_cls.DDP_FIND_UNUSED_PARAMETERS) self.assertIn("MageFlowTransformerBlock", model_cls.MODEL_CLASS._no_split_modules) + self.assertEqual(model_cls.HUGGINGFACE_PATHS["base"], "natalie5/Mage-Flow-Base") def test_context_parallel_is_rejected(self): model_cls = _mageflow_class() @@ -383,6 +385,40 @@ def checkpoint_func(function, *args, **kwargs): self.assertEqual(len(checkpoint_calls), 1) self.assertEqual(checkpoint_calls[0], {"use_reentrant": False}) + def test_model_level_gradient_checkpointing_reaches_wrapped_transformer(self): + model_cls = _mageflow_class() + foundation = object.__new__(model_cls) + transformer = _tiny_transformer() + foundation.model = SimpleNamespace(base_model=SimpleNamespace(model=transformer)) + foundation.unwrap_model = lambda model=None, keep_fp32_wrapper=True: model + + foundation.enable_gradient_checkpointing() + self.assertTrue(transformer.checkpoint) + self.assertTrue(transformer.config.checkpoint) + + foundation.disable_gradient_checkpointing() + self.assertFalse(transformer.checkpoint) + self.assertFalse(transformer.config.checkpoint) + + def test_missing_pipeline_scheduler_config_uses_default_flow_scheduler(self): + with tempfile.TemporaryDirectory() as repo_dir: + scheduler = _load_scheduler(repo_dir) + + self.assertEqual(scheduler.config.num_train_timesteps, 1000) + self.assertEqual(scheduler.config.shift, 6.0) + + def test_training_noise_schedule_uses_default_when_checkpoint_has_no_scheduler_config(self): + model_cls = _mageflow_class() + foundation = object.__new__(model_cls) + foundation.config = SimpleNamespace(flow_schedule_shift=7.0) + + with patch.object(ImageModelFoundation, "setup_training_noise_schedule", side_effect=OSError("missing")): + config, scheduler = foundation.setup_training_noise_schedule() + + self.assertIs(config, foundation.config) + self.assertEqual(scheduler.config.num_train_timesteps, 1000) + self.assertEqual(scheduler.config.shift, 7.0) + def test_transformer_supports_twinflow_time_sign(self): model = _tiny_transformer(enable_time_sign_embed=True) diff --git a/tests/test_metadata_backend.py b/tests/test_metadata_backend.py index 8494bd60c..656604ec1 100644 --- a/tests/test_metadata_backend.py +++ b/tests/test_metadata_backend.py @@ -345,9 +345,6 @@ def test_eval_dataset_single_image_no_validation_error(self): mock_accelerator = MagicMock() mock_accelerator.num_processes = 1 mock_accelerator.process_index = 0 - mock_accelerator.split_between_processes = MagicMock() - mock_accelerator.split_between_processes.return_value.__enter__ = MagicMock(return_value=["eval_img1.jpg"]) - mock_accelerator.split_between_processes.return_value.__exit__ = MagicMock(return_value=False) mock_backend.accelerator = mock_accelerator with ( @@ -416,9 +413,6 @@ def test_eval_dataset_ignores_gradient_accumulation(self): mock_accelerator = MagicMock() mock_accelerator.num_processes = 1 mock_accelerator.process_index = 0 - mock_accelerator.split_between_processes = MagicMock() - mock_accelerator.split_between_processes.return_value.__enter__ = MagicMock(return_value=["eval_img1.jpg"]) - mock_accelerator.split_between_processes.return_value.__exit__ = MagicMock(return_value=False) mock_backend.accelerator = mock_accelerator with ( @@ -440,6 +434,306 @@ def test_eval_dataset_ignores_gradient_accumulation(self): self.assertIn("1.0", mock_backend.aspect_ratio_bucket_indices) +class TestDistributedBucketPadding(unittest.TestCase): + """Regression coverage for PR #2897 distributed bucket policies.""" + + @staticmethod + def _backend(images, *, num_processes=8): + backend = MagicMock(spec=MetadataBackend) + backend.id = "distributed-padding" + backend.batch_size = 1 + backend.repeats = 0 + backend.bucket_report = None + backend.dataset_type = DatasetType.IMAGE + backend.aspect_ratio_bucket_indices = {"1.0": list(images)} + backend.read_only = False + backend.accelerator = MagicMock( + num_processes=num_processes, + process_index=0, + is_main_process=True, + ) + return backend + + def _split_rank(self, images, rank, *, cp_size, apply_padding, repeats=0): + backend = self._backend(images) + backend.repeats = repeats + with ( + patch.dict("os.environ", {"SIMPLETUNER_SHUFFLE_BUCKETS": "0"}), + patch.object( + StateTracker, + "get_args", + return_value=SimpleNamespace(allow_dataset_oversubscription=apply_padding), + ), + patch.object( + StateTracker, + "get_data_backend_config", + return_value={"repeats": repeats} if repeats else {}, + ), + patch( + "simpletuner.helpers.metadata.backends.base.get_cp_aware_dp_info", + return_value=(8, rank, cp_size), + ), + ): + MetadataBackend.split_buckets_between_processes( + backend, + gradient_accumulation_steps=1, + apply_padding=apply_padding, + ) + return backend.aspect_ratio_bucket_indices["1.0"] + + def test_cp_padding_fills_empty_dp_shards_from_final_global_item(self): + images = ["img0", "img1", "img2", "img3"] + shards = [ + self._split_rank(images, rank, cp_size=2, apply_padding=True, repeats=1) + for rank in range(8) + ] + + self.assertEqual([len(shard) for shard in shards], [1] * 8) + self.assertEqual([item for shard in shards for item in shard], images + ["img3"] * 4) + + def test_cp_no_padding_uses_balanced_qr_partition(self): + images = [f"img{index}" for index in range(9)] + shards = [self._split_rank(images, rank, cp_size=2, apply_padding=False) for rank in range(8)] + + self.assertEqual([len(shard) for shard in shards], [2, 1, 1, 1, 1, 1, 1, 1]) + self.assertEqual([item for shard in shards for item in shard], images) + + def test_standard_split_preserves_balanced_qr_order(self): + images = [f"img{index}" for index in range(9)] + shards = [self._split_rank(images, rank, cp_size=1, apply_padding=False) for rank in range(8)] + + self.assertEqual([len(shard) for shard in shards], [2, 1, 1, 1, 1, 1, 1, 1]) + self.assertEqual([item for shard in shards for item in shard], images) + + def test_standard_split_padding_exact_division_preserves_cardinality(self): + images = [f"img{index}" for index in range(8)] + shards = [self._split_rank(images, rank, cp_size=1, apply_padding=True) for rank in range(8)] + + self.assertEqual([len(shard) for shard in shards], [1] * 8) + self.assertEqual([item for shard in shards for item in shard], images) + + def test_padding_non_divisible_preserves_qr_boundaries(self): + images = [f"img{index}" for index in range(10)] + expected = [["img0", "img1"], ["img2", "img3"]] + [[f"img{index}", "img9"] for index in range(4, 10)] + + for cp_size in (1, 2): + with self.subTest(cp_size=cp_size): + shards = [self._split_rank(images, rank, cp_size=cp_size, apply_padding=True) for rank in range(8)] + self.assertEqual(shards, expected) + + +class TestAutomaticOversubscriptionLogicalSequence(unittest.TestCase): + """Automatic repeats are materialised into a rank-local cyclic schedule.""" + + @staticmethod + def _backend(images, dp_size, rank, cp_size=1, batch_size=1, repeats=0): + backend = object.__new__(MetadataBackend) + backend.id = "auto-oversubscription" + backend.batch_size = batch_size + backend.repeats = repeats + backend.bucket_report = None + backend.dataset_type = DatasetType.IMAGE + backend.read_only = False + backend._aspect_ratio_bucket_indices = {"1.0": list(images)} + backend.accelerator = SimpleNamespace( + num_processes=dp_size * cp_size, + process_index=rank, + is_main_process=True, + ) + return backend + + def _split(self, backend, dp_size, rank, cp_size, gradient_accumulation_steps, config=None): + with ( + patch.dict("os.environ", {"SIMPLETUNER_SHUFFLE_BUCKETS": "0"}), + patch.object( + StateTracker, + "get_args", + return_value=SimpleNamespace(allow_dataset_oversubscription=True, seed=0), + ), + patch.object(StateTracker, "get_data_backend_config", return_value=config or {}), + patch( + "simpletuner.helpers.metadata.backends.base.get_cp_aware_dp_info", + return_value=(dp_size, rank, cp_size), + ), + ): + MetadataBackend.split_buckets_between_processes( + backend, + gradient_accumulation_steps=gradient_accumulation_steps, + apply_padding=True, + ) + + def test_standard_auto_repeat_provides_one_complete_ga_window_per_rank(self): + shards = [] + for rank in range(8): + backend = self._backend(["img0"], dp_size=8, rank=rank) + self._split(backend, 8, rank, 1, gradient_accumulation_steps=4, config={"repeats": 0}) + shards.append(backend.aspect_ratio_bucket_indices["1.0"]) + self.assertEqual(backend.repeats, 0) + + self.assertEqual([len(shard) for shard in shards], [4] * 8) + self.assertEqual(sum(map(len, shards)), 32) + self.assertEqual({item for shard in shards for item in shard}, {"img0"}) + + def test_context_parallel_auto_repeat_uses_effective_dp_and_cyclic_tail(self): + images = [f"img{index}" for index in range(5)] + shards = [] + for rank in range(4): + backend = self._backend(images, dp_size=4, rank=rank, cp_size=2) + self._split(backend, 4, rank, 2, gradient_accumulation_steps=2, config={"repeats": 0}) + shards.append(backend.aspect_ratio_bucket_indices["1.0"]) + + self.assertEqual([len(shard) for shard in shards], [4] * 4) + self.assertEqual( + [item for shard in shards for item in shard], + images + images + (images * 2)[:6], + ) + + def test_docs_style_auto_repeat_cardinality_is_padded_to_full_window(self): + images = [f"img{index}" for index in range(25)] + shards = [] + for rank in range(8): + backend = self._backend(images, dp_size=8, rank=rank) + self._split(backend, 8, rank, 1, gradient_accumulation_steps=4, config={"repeats": 0}) + shards.append(backend.aspect_ratio_bucket_indices["1.0"]) + + self.assertEqual([len(shard) for shard in shards], [8] * 8) + self.assertEqual( + [item for shard in shards for item in shard], + images + images + images[:14], + ) + + def test_auto_repeat_count_is_backend_wide_across_buckets(self): + backend = self._backend([], dp_size=2, rank=0, batch_size=4) + backend._aspect_ratio_bucket_indices = { + "small": ["s0", "s1", "s2"], + "large": ["l0", "l1", "l2", "l3", "l4", "l5", "l6"], + } + self._split(backend, 2, 0, 1, gradient_accumulation_steps=1, config={"repeats": 0}) + + self.assertEqual(backend.repeats, 0) + self.assertEqual( + {bucket: len(values) for bucket, values in backend.aspect_ratio_bucket_indices.items()}, + {"small": 8, "large": 12}, + ) + self.assertEqual( + backend.aspect_ratio_bucket_indices["small"], + ["s0", "s1", "s2", "s0", "s1", "s2", "s0", "s1"], + ) + self.assertEqual( + backend.aspect_ratio_bucket_indices["large"], + ["l0", "l1", "l2", "l3", "l4", "l5", "l6", "l0", "l1", "l2", "l3", "l4"], + ) + + def test_manual_repeats_are_not_materialised(self): + backend = self._backend(["img0"], dp_size=8, rank=0, repeats=31) + self._split(backend, 8, 0, 1, gradient_accumulation_steps=4, config={"repeats": 31}) + + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], ["img0"]) + self.assertEqual(backend.repeats, 31) + + def test_empty_buckets_are_ignored_without_division_by_zero(self): + backend = self._backend([], dp_size=8, rank=0, batch_size=1) + backend._aspect_ratio_bucket_indices = {"empty": [], "full": ["img0"]} + self._split(backend, 8, 0, 1, gradient_accumulation_steps=1, config={"repeats": 0}) + self.assertEqual(backend.aspect_ratio_bucket_indices["empty"], []) + self.assertEqual(len(backend.aspect_ratio_bucket_indices["full"]), 1) + + +class TestEmptyBucketRepeatValidation(unittest.TestCase): + """An empty bucket key survives update_buckets_with_existing_files and reaches the split.""" + + @staticmethod + def _backend(buckets, *, repeats=0): + backend = MagicMock(spec=MetadataBackend) + backend.id = "empty-bucket" + backend.batch_size = 1 + backend.repeats = repeats + backend.bucket_report = None + backend.dataset_type = DatasetType.IMAGE + backend.aspect_ratio_bucket_indices = {key: list(value) for key, value in buckets.items()} + backend.read_only = False + backend.accelerator = MagicMock(num_processes=1, process_index=0, is_main_process=True) + return backend + + @staticmethod + def _split(backend, *, num_processes, allow_oversubscription, dp_rank=0): + with ( + patch.dict("os.environ", {"SIMPLETUNER_SHUFFLE_BUCKETS": "0"}), + patch.object( + StateTracker, + "get_args", + return_value=SimpleNamespace(allow_dataset_oversubscription=allow_oversubscription), + ), + patch.object(StateTracker, "get_data_backend_config", return_value={}), + patch( + "simpletuner.helpers.metadata.backends.base.get_cp_aware_dp_info", + return_value=(num_processes, dp_rank, 1), + ), + ): + MetadataBackend.split_buckets_between_processes( + backend, + gradient_accumulation_steps=1, + apply_padding=allow_oversubscription, + ) + + def test_empty_bucket_does_not_divide_by_zero(self): + # The crash is not distribution-specific: effective_batch_size is at least 1, so an + # empty bucket always joins buckets_that_will_fail regardless of world size. + for num_processes in (1, 2, 8): + for allow_oversubscription in (True, False): + with self.subTest(num_processes=num_processes, allow_oversubscription=allow_oversubscription): + backend = self._backend({"1.0": ["a.jpg", "b.jpg"], "1.5": []}) + try: + self._split(backend, num_processes=num_processes, allow_oversubscription=allow_oversubscription) + except ZeroDivisionError as error: + self.fail(f"empty bucket divided by zero: {error}") + except ValueError: + # The pre-existing "zero usable batches" guard is allowed to fire; it is + # asserted on its own below. + pass + + def test_oversubscription_ignores_the_empty_bucket_when_adjusting_repeats(self): + # The auto repeat factor ceil(8 / 2) - 1 = 3 is driven by the two-image bucket alone. + # Since #2918, repeats is left unmutated and the factor is materialised into the + # rank-local shard instead: logical 2 * (3 + 1) = 8 samples, batch-aligned to the + # effective batch size of 8, so every rank holds exactly 8 // 8 = 1 cyclic sample. + shards = [] + for dp_rank in range(8): + backend = self._backend({"1.0": ["a.jpg", "b.jpg"], "1.5": []}) + self._split(backend, num_processes=8, allow_oversubscription=True, dp_rank=dp_rank) + self.assertEqual(backend.repeats, 0, "auto oversubscription must not mutate repeats") + self.assertEqual(backend.aspect_ratio_bucket_indices["1.5"], []) + shards.append(backend.aspect_ratio_bucket_indices["1.0"]) + + self.assertEqual([len(shard) for shard in shards], [1] * 8) + self.assertEqual([shard[0] for shard in shards], ["a.jpg", "b.jpg"] * 4) + + def test_empty_bucket_still_reports_zero_usable_batches_without_oversubscription(self): + backend = self._backend({"1.0": ["a.jpg", "b.jpg"], "1.5": []}) + with self.assertRaises(ValueError) as raised: + self._split(backend, num_processes=8, allow_oversubscription=False) + self.assertIn("zero usable batches", str(raised.exception)) + + def test_bucket_without_empty_keys_is_unaffected(self): + backend = self._backend({"1.0": ["a.jpg", "b.jpg"]}) + self._split(backend, num_processes=8, allow_oversubscription=True) + + # Same materialised schedule as the empty-bucket case: repeats stays 0 and rank 0 + # holds the first of the eight batch-aligned cyclic samples. + self.assertEqual(backend.repeats, 0, "auto oversubscription must not mutate repeats") + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], ["a.jpg"]) + + def test_refresh_leaves_an_empty_bucket_key_behind(self): + # update_buckets_with_existing_files assigns [] rather than dropping the key, which is + # how the empty bucket reaches the split in the first place. + backend = MagicMock() + backend.aspect_ratio_bucket_indices = {"1.0": ["a.jpg"], "1.5": ["gone.jpg"]} + backend.bucket_report = None + MetadataBackend.update_buckets_with_existing_files(backend, {"a.jpg"}) + + self.assertEqual(backend.aspect_ratio_bucket_indices, {"1.0": ["a.jpg"], "1.5": []}) + + class TestFilteringStatistics(unittest.TestCase): """Test filtering_statistics storage and retrieval in metadata backends (issue #2474).""" diff --git a/tests/test_musubi_block_swap.py b/tests/test_musubi_block_swap.py index a6e112173..bff4108e3 100644 --- a/tests/test_musubi_block_swap.py +++ b/tests/test_musubi_block_swap.py @@ -51,6 +51,40 @@ def test_quanto_qlinear_streams_without_apply_swap(self): self.assertEqual(qlinear.weight._data.device.type, "cpu") self.assertEqual(qlinear.weight._scale.device.type, "cpu") + def test_sdnq_module_streams_without_apply_swap(self): + device = self._accelerator_device() + + class FakeSDNQLinear(nn.Linear): + pass + + FakeSDNQLinear.__module__ = "sdnq.training.layers.linear" + block = nn.Sequential(FakeSDNQLinear(4, 4), nn.SiLU(), FakeSDNQLinear(4, 4)) + for param in block.parameters(): + param.requires_grad_(False) + param.sdnq_dequantizer = object() + param.weight = param.detach().clone() + param.scale = torch.ones(param.shape[0], 1, device=param.device) + block.to(device) + + manager = MusubiBlockSwapManager( + block_indices=[0], + offload_device=torch.device("cpu"), + logger=logging.getLogger(__name__), + ) + + with patch.object(FakeSDNQLinear, "_apply", side_effect=RuntimeError("_apply(): Couldn't swap SDNQLinear.weight")): + manager.stream_out(block) + self.assertTrue(_module_on_device(block, torch.device("cpu"))) + self.assertEqual(block[0].weight.device.type, "cpu") + self.assertEqual(block[0].weight.scale.device.type, "cpu") + + manager.stream_in(block, device) + self.assertTrue(_module_on_device(block, device)) + self.assertEqual(block[0].weight.device.type, device.type) + self.assertEqual(block[0].weight.scale.device.type, device.type) + output = block(torch.randn(2, 4, device=device)) + self.assertEqual(output.device.type, device.type) + def test_stream_out_keeps_trainable_params_on_accelerator(self): device = self._accelerator_device() block = nn.Sequential(nn.Linear(4, 4), nn.SiLU(), nn.Linear(4, 4)) diff --git a/tests/test_omnigen_model.py b/tests/test_omnigen_model.py index a249930f6..9b06ff2a1 100644 --- a/tests/test_omnigen_model.py +++ b/tests/test_omnigen_model.py @@ -35,6 +35,9 @@ def setUp(self): def test_model_supports_crepa_self_flow(self): self.assertTrue(self.model.supports_crepa_self_flow()) + def test_model_does_not_use_text_embedding_cache(self): + self.assertFalse(self.model.uses_text_embeddings_cache()) + def test_prepare_crepa_self_flow_batch_creates_tokenwise_timesteps(self): self.model.config.crepa_self_flow_mask_ratio = 0.5 batch = { diff --git a/tests/test_packed_attention_processors.py b/tests/test_packed_attention_processors.py index 055031425..ac3b6212d 100644 --- a/tests/test_packed_attention_processors.py +++ b/tests/test_packed_attention_processors.py @@ -1,4 +1,5 @@ import unittest +from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import patch @@ -326,6 +327,52 @@ def test_flux2_single_stream_context_parallel_uses_distributed_dispatch(self): self.assertIs(dispatch.call_args.kwargs["parallel_config"], parallel_config) self.assertEqual(dispatch.call_args.kwargs["attn_mask"].shape, (1, 1, 1, 6)) + def test_flux2_metal_flash_fast_path_uses_activation_offload_context(self): + context_calls = [] + + @contextmanager + def fake_activation_offload_context(enabled, label=None): + context_calls.append((enabled, label)) + yield + + def fake_metal_attention(query, *_args, **_kwargs): + return torch.zeros_like(query) + + double_stream = Flux2Attention(query_dim=8, heads=2, dim_head=4, out_dim=8) + single_stream = Flux2ParallelSelfAttention(query_dim=8, heads=2, dim_head=4, out_dim=8, mlp_ratio=1.0) + + with ( + patch( + "simpletuner.helpers.models.flux2.transformer.activation_offload_context", + new=fake_activation_offload_context, + ), + patch( + "simpletuner.helpers.models.flux2.transformer.maybe_metal_flash_rope_attention", + side_effect=fake_metal_attention, + ) as metal_attention, + ): + double_output = double_stream( + torch.randn(1, 2, 8), + image_rotary_emb=object(), + offload_attention=True, + ) + single_output = single_stream( + torch.randn(1, 3, 8), + image_rotary_emb=object(), + offload_attention=True, + ) + + self.assertEqual(double_output.shape, (1, 2, 8)) + self.assertEqual(single_output.shape, (1, 3, 8)) + self.assertEqual(metal_attention.call_count, 2) + self.assertEqual( + context_calls, + [ + (True, "Flux2Attention:attention"), + (True, "Flux2ParallelSelfAttention:attention"), + ], + ) + def test_ltx2_transformer_fuse_enables_packed_self_attention_processors(self): model = LTX2VideoTransformer3DModel( in_channels=4, diff --git a/tests/test_pipelines/test_wan_s2v_pipeline.py b/tests/test_pipelines/test_wan_s2v_pipeline.py index 1920201ea..85c67316c 100644 --- a/tests/test_pipelines/test_wan_s2v_pipeline.py +++ b/tests/test_pipelines/test_wan_s2v_pipeline.py @@ -1,5 +1,7 @@ """Tests for WanS2V (Speech-to-Video) pipeline components.""" +import os +import tempfile import unittest from unittest import mock @@ -84,6 +86,24 @@ def test_model_huggingface_path(self): "tolgacangoz/Wan2.2-S2V-14B-Diffusers", ) + def test_checkpoint_tensor_loader_returns_none_for_missing_or_unreadable_index(self): + from simpletuner.helpers.models.wan_s2v.model import WanS2V + + model = WanS2V.__new__(WanS2V) + with tempfile.TemporaryDirectory() as tempdir: + self.assertIsNone(model._load_checkpoint_tensor(tempdir, "transformer", "missing.weight", {})) + + transformer_dir = os.path.join(tempdir, "transformer") + os.makedirs(transformer_dir, exist_ok=True) + with open( + os.path.join(transformer_dir, "diffusion_pytorch_model.safetensors.index.json"), + "w", + encoding="utf-8", + ) as index_file: + index_file.write("{") + + self.assertIsNone(model._load_checkpoint_tensor(tempdir, "transformer", "missing.weight", {})) + class TestWanS2VModelMetadata(unittest.TestCase): """Test that WanS2V is registered in model metadata.""" diff --git a/tests/test_qwen_text_encoder_override.py b/tests/test_qwen_text_encoder_override.py new file mode 100644 index 000000000..1e4e423fc --- /dev/null +++ b/tests/test_qwen_text_encoder_override.py @@ -0,0 +1,155 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +from simpletuner.helpers.models.common import ImageModelFoundation + + +class QwenTextModel: + pass + + +class QwenTokenizer: + pass + + +class ClipTextModel: + pass + + +class ClipTokenizer: + pass + + +class DummyQwenFoundation(ImageModelFoundation): + NAME = "Dummy Qwen" + + def model_predict(self, prepared_batch, custom_timesteps: list = None): + raise NotImplementedError + + def _encode_prompts(self, prompts: list, is_negative_prompt: bool = False): + raise NotImplementedError + + def convert_text_embed_for_pipeline(self, text_embedding): + raise NotImplementedError + + def convert_negative_text_embed_for_pipeline(self, text_embedding): + raise NotImplementedError + + +class QwenTextEncoderOverrideTests(unittest.TestCase): + def _model(self, text_encoder_configuration): + model = object.__new__(DummyQwenFoundation) + model.config = SimpleNamespace( + model_family="qwen_image", + pretrained_model_name_or_path="base/model", + qwen_text_encoder_model_name_or_path="custom/qwen", + ) + model.TEXT_ENCODER_CONFIGURATION = text_encoder_configuration + return model + + def test_single_qwen_encoder_uses_qwen_override_and_clears_component_subfolders(self): + qwen_config = { + "name": "Qwen2.5-VL", + "tokenizer": QwenTokenizer, + "tokenizer_subfolder": "tokenizer", + "model": QwenTextModel, + "subfolder": "text_encoder", + } + clip_config = { + "name": "CLIP-L/14", + "tokenizer": ClipTokenizer, + "tokenizer_subfolder": "tokenizer_2", + "model": ClipTextModel, + "subfolder": "text_encoder_2", + } + model = self._model( + { + "text_encoder": qwen_config, + "text_encoder_2": clip_config, + } + ) + + self.assertEqual(model._resolve_text_encoder_path(qwen_config), "custom/qwen") + self.assertIsNone(model._resolve_text_encoder_subfolder(qwen_config, "subfolder", "text_encoder")) + self.assertIsNone(model._resolve_text_encoder_subfolder(qwen_config, "tokenizer_subfolder", "tokenizer")) + + self.assertEqual(model._resolve_text_encoder_path(clip_config), "base/model") + self.assertEqual( + model._resolve_text_encoder_subfolder(clip_config, "subfolder", "text_encoder"), + "text_encoder_2", + ) + + def test_multiple_qwen_encoders_ignore_override_and_warn(self): + first_qwen_config = { + "name": "Qwen3-A", + "tokenizer": QwenTokenizer, + "model": QwenTextModel, + "subfolder": "text_encoder", + } + second_qwen_config = { + "name": "Qwen3-B", + "tokenizer": QwenTokenizer, + "model": QwenTextModel, + "subfolder": "text_encoder_2", + } + model = self._model( + { + "text_encoder": first_qwen_config, + "text_encoder_2": second_qwen_config, + } + ) + + with patch("simpletuner.helpers.models.common.logger.warning") as warning: + self.assertEqual(model._resolve_text_encoder_path(first_qwen_config), "base/model") + + warning.assert_called_once() + self.assertIn("Ignoring qwen_text_encoder_model_name_or_path", warning.call_args.args[0]) + self.assertEqual(warning.call_args.args[2], 2) + self.assertEqual( + model._resolve_text_encoder_subfolder(first_qwen_config, "subfolder", "text_encoder"), + "text_encoder", + ) + + def test_webui_field_and_cli_parser_include_qwen_override(self): + from simpletuner.helpers.configuration.cmd_args import get_argument_parser + from simpletuner.simpletuner_sdk.server.services.field_registry.registry import FieldRegistry + + registry = FieldRegistry() + field = registry.get_field("qwen_text_encoder_model_name_or_path") + + self.assertIsNotNone(field) + self.assertEqual(field.arg_name, "--qwen_text_encoder_model_name_or_path") + self.assertEqual(field.documentation, "OPTIONS.md#--qwen_text_encoder_model_name_or_path") + + parser = get_argument_parser() + args = parser.parse_args( + [ + "--model_family", + "krea2", + "--output_dir", + "/tmp/simpletuner-test", + "--model_type", + "lora", + "--optimizer", + "adamw_bf16", + "--data_backend_config", + "/tmp/backend.json", + "--qwen_text_encoder_model_name_or_path", + "custom/qwen", + ] + ) + + self.assertEqual(args.qwen_text_encoder_model_name_or_path, "custom/qwen") + + def test_env_mapping_includes_qwen_override(self): + from simpletuner.helpers.configuration.env_file import env_to_args_map + + self.assertEqual( + env_to_args_map["QWEN_TEXT_ENCODER_MODEL_NAME_OR_PATH"], + "--qwen_text_encoder_model_name_or_path", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_ramtorch.py b/tests/test_ramtorch.py index 2b457ab26..9ec6db2fb 100644 --- a/tests/test_ramtorch.py +++ b/tests/test_ramtorch.py @@ -192,6 +192,20 @@ def __init__(self): self.assertEqual(model.ramtorch_buffer.device.type, "cpu") self.assertTrue(getattr(model.ramtorch_buffer, "is_ramtorch", False)) + def test_ramtorch_profile_reset_resets_cuda_peak_memory_stats(self): + from simpletuner.helpers.ramtorch import profiling as ramtorch_profile + + with ( + patch("torch.cuda.is_available", return_value=True), + patch("torch.cuda.device_count", return_value=2), + patch("torch.cuda.reset_peak_memory_stats") as reset_peak_memory_stats, + ): + ramtorch_profile.reset_for_new_run() + + self.assertEqual(reset_peak_memory_stats.call_count, 2) + self.assertEqual(str(reset_peak_memory_stats.call_args_list[0].args[0]), "cuda:0") + self.assertEqual(str(reset_peak_memory_stats.call_args_list[1].args[0]), "cuda:1") + def test_torchao_int8_ramtorch_backward_uses_dense_weight_view(self): try: from torchao.prototype.quantized_training import int8_weight_only_quantized_training @@ -313,9 +327,18 @@ def __init__(self): ignored = ramtorch_utils.mark_ddp_ignore_params(model) self.assertEqual(ignored, 3) ignore_set = getattr(model, "_ddp_params_and_buffers_to_ignore", set()) + self.assertNotIn("linear_a.weight", ignore_set) + self.assertNotIn("linear_a.bias", ignore_set) self.assertIn("linear_b.weight", ignore_set) self.assertIn("linear_b.bias", ignore_set) self.assertIn("ramtorch_buffer", ignore_set) + self.assertEqual(ramtorch_utils.mark_ddp_ignore_params(model), 0) + + def test_mark_ddp_ignore_params_skips_plain_cpu_modules(self): + model = nn.Sequential(nn.Linear(2, 2), nn.LayerNorm(2)) + ignored = ramtorch_utils.mark_ddp_ignore_params(model) + self.assertEqual(ignored, 0) + self.assertFalse(hasattr(model, "_ddp_params_and_buffers_to_ignore")) def test_prefetch_hooks_follow_ramtorch_module_order(self): from simpletuner.helpers.ramtorch_extensions import add_ramtorch_prefetch_hooks @@ -622,6 +645,53 @@ def forward(self, x): with patch.dict("os.environ", {"SIMPLETUNER_RAMTORCH_PREFETCH_POLICY": "sync"}): self.assertEqual(add_ramtorch_prefetch_hooks(model), []) + def test_prefetch_hooks_start_successor_before_current_forward(self): + from simpletuner.helpers.ramtorch_extensions import add_ramtorch_prefetch_hooks + from simpletuner.helpers.training.state_tracker import StateTracker + + StateTracker.reset_ramtorch_prefetch_orders() + + class _PrefetchModule(nn.Module): + is_ramtorch = True + + def __init__(self, label, calls): + super().__init__() + self.label = label + self.calls = calls + + def prefetch_forward(self): + self.calls.append(("prefetch", self.label)) + return True + + def forward(self, x): + self.calls.append(("forward", self.label)) + return x + + calls = [] + model = nn.Sequential( + _PrefetchModule("first", calls), + _PrefetchModule("second", calls), + _PrefetchModule("third", calls), + ) + + hooks = add_ramtorch_prefetch_hooks(model, component_label="prehook-timing-test") + try: + model(torch.ones(1)) + finally: + for hook in hooks: + hook.remove() + + self.assertEqual( + calls, + [ + ("prefetch", "second"), + ("forward", "first"), + ("prefetch", "third"), + ("forward", "second"), + ("forward", "third"), + ], + ) + def test_bundled_linear_skips_frozen_weight_gradients(self): from simpletuner.helpers.ramtorch.modules.linear import Linear diff --git a/tests/test_sampler.py b/tests/test_sampler.py index 9511452ca..d09cdc018 100644 --- a/tests/test_sampler.py +++ b/tests/test_sampler.py @@ -1,7 +1,9 @@ import logging import os +import tempfile import unittest from math import ceil +from types import SimpleNamespace # Import test configuration to suppress logging/warnings try: @@ -20,16 +22,19 @@ from accelerate import PartialState from PIL import Image +from simpletuner.helpers.data_backend.dataset_types import DatasetType +from simpletuner.helpers.metadata.backends.base import MetadataBackend from simpletuner.helpers.metadata.backends.discovery import DiscoveryMetadataBackend from simpletuner.helpers.multiaspect.sampler import MultiAspectSampler from simpletuner.helpers.multiaspect.state import BucketStateManager +from simpletuner.helpers.training.state_tracker import StateTracker from tests.helpers.data import MockDataBackend class TestMultiAspectSampler(unittest.TestCase): def setUp(self): self.process_state = PartialState() - self.accelerator = MagicMock() + self.accelerator = MagicMock(num_processes=1, process_index=0) self.accelerator.log = MagicMock() self.metadata_backend = Mock(spec=DiscoveryMetadataBackend) self.metadata_backend.id = "foo" @@ -68,6 +73,77 @@ def test_load_buckets(self): buckets = self.sampler.load_buckets() self.assertEqual(buckets, ["1.0"]) + def test_padded_occurrences_are_consumed_individually(self): + """A repeated filepath represents multiple scheduled samples, not one boolean item.""" + metadata_backend = object.__new__(DiscoveryMetadataBackend) + metadata_backend.instance_data_dir = "" + metadata_backend.aspect_ratio_bucket_indices = {"1.0": ["same.jpg", "same.jpg"]} + metadata_backend.seen_images = {} + + sampler = object.__new__(MultiAspectSampler) + sampler.metadata_backend = metadata_backend + sampler.logger = MagicMock() + sampler.debug_log = MagicMock() + + self.assertEqual(sampler._get_unseen_images("1.0"), ["same.jpg", "same.jpg"]) + metadata_backend.mark_as_seen("same.jpg") + self.assertEqual(sampler._get_unseen_images("1.0"), ["same.jpg"]) + + # Old checkpoints stored booleans. True means the filepath was + # exhausted, including every scheduled duplicate. + metadata_backend.seen_images = {"same.jpg": True} + self.assertEqual(sampler._get_unseen_images("1.0"), []) + + metadata_backend.seen_images = {} + metadata_backend.mark_batch_as_seen(["same.jpg", "same.jpg"]) + self.assertEqual(metadata_backend.seen_occurrence_count("same.jpg"), 2) + self.assertEqual(sampler._get_unseen_images("1.0"), []) + + metadata_backend.seen_images = {"same.jpg": "corrupt"} + with self.assertRaisesRegex(TypeError, "Invalid seen occurrence count"): + metadata_backend.seen_occurrence_count("same.jpg") + + def test_load_states_restores_schedule_before_normalizing_legacy_seen_flags(self): + self.sampler.state_manager.load_state.return_value = { + "aspect_ratio_bucket_indices": {"1.0": ["same.jpg", "same.jpg", "other.jpg"]}, + "buckets": ["1.0"], + "current_bucket": 0, + "exhausted_buckets": ["old"], + "seen_images": {"same.jpg": True, "other.jpg": False, "legacy.jpg": True}, + "dp_size": 1, + "dp_rank": 0, + } + + self.metadata_backend.aspect_ratio_bucket_indices = {"1.0": ["stale.jpg"]} + self.metadata_backend.seen_images = {} + self.sampler.load_states(self.state_path) + + self.assertEqual( + self.metadata_backend.aspect_ratio_bucket_indices, + {"1.0": ["same.jpg", "same.jpg", "other.jpg"]}, + ) + self.assertEqual(self.sampler.buckets, ["1.0"]) + self.assertEqual(self.sampler.current_bucket, 0) + self.assertEqual(self.sampler.exhausted_buckets, ["old"]) + self.assertEqual(self.metadata_backend.seen_images["same.jpg"], 2) + self.assertEqual(self.metadata_backend.seen_images["other.jpg"], 0) + self.assertTrue(self.metadata_backend.seen_images["legacy.jpg"]) + + def test_load_states_keeps_the_fresh_split_when_the_checkpoint_records_no_layout(self): + # Checkpoints written before the layout was recorded cannot be attributed to a rank, so + # the schedule is left alone. Seen state still loads. + self.sampler.state_manager.load_state.return_value = { + "aspect_ratio_bucket_indices": {"1.0": ["same.jpg", "same.jpg", "other.jpg"]}, + "seen_images": {"same.jpg": True}, + } + + self.metadata_backend.aspect_ratio_bucket_indices = {"1.0": ["fresh.jpg"]} + self.metadata_backend.seen_images = {} + self.sampler.load_states(self.state_path) + + self.assertEqual(self.metadata_backend.aspect_ratio_bucket_indices, {"1.0": ["fresh.jpg"]}) + self.assertTrue(self.metadata_backend.seen_images["same.jpg"]) + def test_change_bucket(self): self.sampler.buckets = ["1.5"] self.sampler.exhausted_buckets = ["1.0"] @@ -181,5 +257,125 @@ def mock_validate_and_yield_images_from_samples(samples, bucket): self.assertIn("/fake/dir/image4.jpg", result_paths) +class TestSamplerResumeSchedule(unittest.TestCase): + """A checkpointed schedule is one rank's shard; restoring it must keep the ranks partitioned.""" + + @staticmethod + def _split(world_size, rank, shuffle_seed, images): + backend = MagicMock(spec=MetadataBackend) + backend.id = "resume" + backend.batch_size = 1 + backend.repeats = 0 + backend.bucket_report = None + backend.dataset_type = DatasetType.IMAGE + backend.aspect_ratio_bucket_indices = {"1.0": list(images)} + backend.read_only = False + backend.seen_images = {} + backend.accelerator = SimpleNamespace( + num_processes=world_size, + process_index=rank, + is_main_process=rank == 0, + ) + with ( + patch.object( + StateTracker, + "get_args", + return_value=SimpleNamespace(allow_dataset_oversubscription=False, seed=shuffle_seed), + ), + patch.object(StateTracker, "get_data_backend_config", return_value={}), + patch( + "simpletuner.helpers.metadata.backends.base.broadcast_object_from_main", + side_effect=lambda value: shuffle_seed, + ), + ): + MetadataBackend.split_buckets_between_processes( + backend, + gradient_accumulation_steps=1, + apply_padding=False, + ) + return backend + + @staticmethod + def _sampler(backend): + sampler = object.__new__(MultiAspectSampler) + sampler.id = backend.id + sampler.metadata_backend = backend + sampler.accelerator = backend.accelerator + sampler.batch_size = 1 + sampler.buckets = list(backend.aspect_ratio_bucket_indices) + sampler.exhausted_buckets = [] + sampler.current_bucket = None + sampler.current_epoch = 1 + sampler.sample_type_strs = "images" + sampler.logger = MagicMock() + sampler._val_master_list = [] + sampler.state_manager = BucketStateManager(backend.id) + return sampler + + @staticmethod + def _state_path(directory, rank): + filename = "training_state.json" if rank == 0 else f"training_state-rank{rank}.json" + return os.path.join(directory, filename) + + def _resume_shards(self, *, images, save_world_size, saving_ranks, resume_world_size): + with tempfile.TemporaryDirectory() as checkpoint_dir: + for rank in range(save_world_size): + backend = self._split(save_world_size, rank, shuffle_seed=1, images=images) + if rank in saving_ranks: + self._sampler(backend).save_state(self._state_path(checkpoint_dir, rank)) + + shards = [] + for rank in range(resume_world_size): + # A relaunch without --seed draws a new shuffle seed, so the fresh split differs + # from the one the checkpoint was written against. + backend = self._split(resume_world_size, rank, shuffle_seed=2, images=images) + self._sampler(backend).load_states(self._state_path(checkpoint_dir, rank)) + shards.append(list(backend.aspect_ratio_bucket_indices["1.0"])) + return shards + + def _assert_partitions(self, shards, images): + scheduled = [path for shard in shards for path in shard] + self.assertEqual( + sorted(scheduled), + sorted(images), + msg=f"ranks no longer partition the dataset: {shards}", + ) + + def test_resume_after_a_rank_zero_only_save_keeps_the_ranks_partitioned(self): + images = [f"image-{index:02d}.jpg" for index in range(12)] + shards = self._resume_shards(images=images, save_world_size=4, saving_ranks={0}, resume_world_size=4) + self._assert_partitions(shards, images) + + def test_resume_when_every_rank_saved_keeps_the_ranks_partitioned(self): + # Control: a complete checkpoint is restorable and stays a partition. + images = [f"image-{index:02d}.jpg" for index in range(12)] + shards = self._resume_shards(images=images, save_world_size=4, saving_ranks=set(range(4)), resume_world_size=4) + self._assert_partitions(shards, images) + + def test_resume_when_no_rank_saved_keeps_the_ranks_partitioned(self): + # Control: with nothing to restore the fresh split is already a partition. Together with + # the previous control this isolates the mixed state as the cause. + images = [f"image-{index:02d}.jpg" for index in range(12)] + shards = self._resume_shards(images=images, save_world_size=4, saving_ranks=set(), resume_world_size=4) + self._assert_partitions(shards, images) + + def test_resume_when_rank_zero_alone_did_not_save_keeps_the_ranks_partitioned(self): + # The completeness check has to cover rank 0 as well: its state file is the one that is + # named differently, so scanning only the -rank{N} siblings would miss it. + images = [f"image-{index:02d}.jpg" for index in range(12)] + shards = self._resume_shards(images=images, save_world_size=4, saving_ranks={1, 2, 3}, resume_world_size=4) + self._assert_partitions(shards, images) + + def test_resume_onto_a_smaller_world_size_keeps_the_ranks_partitioned(self): + images = [f"image-{index:02d}.jpg" for index in range(16)] + shards = self._resume_shards(images=images, save_world_size=8, saving_ranks={0}, resume_world_size=4) + self._assert_partitions(shards, images) + + def test_resume_onto_a_larger_world_size_keeps_the_ranks_partitioned(self): + images = [f"image-{index:02d}.jpg" for index in range(16)] + shards = self._resume_shards(images=images, save_world_size=4, saving_ranks=set(range(4)), resume_world_size=8) + self._assert_partitions(shards, images) + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_save_hooks.py b/tests/test_save_hooks.py new file mode 100644 index 000000000..4de40c166 --- /dev/null +++ b/tests/test_save_hooks.py @@ -0,0 +1,129 @@ +import json +import multiprocessing +import tempfile +import unittest +from contextlib import nullcontext +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from accelerate.utils import DistributedType + +from simpletuner.helpers.training.save_hooks import SaveHookManager +from simpletuner.helpers.training.state_tracker import StateTracker + + +def _concurrent_writer(directory, barrier, iterations, errors): + from simpletuner.helpers.training.state_tracker import StateTracker as tracker + + for _ in range(iterations): + barrier.wait() + try: + tracker.save_ramtorch_prefetch_orders(directory) + except FileNotFoundError: + with errors.get_lock(): + errors.value += 1 + + +class SaveHookManagerTests(unittest.TestCase): + def _run_save_model_hook(self, is_local_main_process): + output_dir = "/tmp/checkpoint" + + with ( + patch("simpletuner.helpers.training.save_hooks.StateTracker.save_training_state") as save_state, + patch("simpletuner.helpers.training.save_hooks.StateTracker.save_ramtorch_prefetch_orders") as save_orders, + ): + manager = object.__new__(SaveHookManager) + manager.accelerator = SimpleNamespace( + distributed_type=DistributedType.NO, + is_main_process=is_local_main_process, + is_local_main_process=is_local_main_process, + ) + manager.args = SimpleNamespace(model_type="full") + manager.training_state_path = "training_state.json" + manager._offload_models_during_save = Mock(return_value=nullcontext()) + manager._save_ema_state = Mock() + manager._is_fsdp2 = Mock(return_value=False) + manager._save_full_model = Mock() + + manager.save_model_hook([], [], output_dir) + + return save_state, save_orders + + def test_ramtorch_prefetch_orders_are_saved_only_on_local_main_process(self): + for is_local_main_process in (True, False): + with self.subTest(is_local_main_process=is_local_main_process): + _, save_orders = self._run_save_model_hook(is_local_main_process) + self.assertEqual(save_orders.call_count, int(is_local_main_process)) + + def test_ramtorch_prefetch_orders_ignore_global_main_flag(self): + # A non-zero node's local main has is_main_process=False but must still + # write, so the file lands on that node's storage in multi-node runs. + output_dir = "/tmp/checkpoint" + + with ( + patch("simpletuner.helpers.training.save_hooks.StateTracker.save_training_state"), + patch("simpletuner.helpers.training.save_hooks.StateTracker.save_ramtorch_prefetch_orders") as save_orders, + ): + manager = object.__new__(SaveHookManager) + manager.accelerator = SimpleNamespace( + distributed_type=DistributedType.NO, + is_main_process=False, + is_local_main_process=True, + ) + manager.args = SimpleNamespace(model_type="full") + manager.training_state_path = "training_state.json" + manager._offload_models_during_save = Mock(return_value=nullcontext()) + manager._save_ema_state = Mock() + manager._is_fsdp2 = Mock(return_value=False) + manager._save_full_model = Mock() + + manager.save_model_hook([], [], output_dir) + + self.assertEqual(save_orders.call_count, 1) + + def test_training_state_is_saved_on_every_rank(self): + for is_local_main_process in (True, False): + with self.subTest(is_local_main_process=is_local_main_process): + save_state, _ = self._run_save_model_hook(is_local_main_process) + self.assertEqual(save_state.call_count, 1) + + +class RamtorchPrefetchOrderWriterTests(unittest.TestCase): + def test_writer_leaves_only_the_final_file_behind(self): + with tempfile.TemporaryDirectory() as tmpdir: + StateTracker.reset_ramtorch_prefetch_orders() + StateTracker.save_ramtorch_prefetch_orders(tmpdir) + + entries = sorted(p.name for p in Path(tmpdir).iterdir()) + self.assertEqual(entries, ["ramtorch_prefetch_orders.json"]) + with (Path(tmpdir) / "ramtorch_prefetch_orders.json").open() as handle: + self.assertEqual(json.load(handle), {"version": 1, "components": {}}) + + def test_concurrent_writers_do_not_race(self): + # Two node-mains on a shared filesystem write the same path at once. + # With a shared temp filename this raises FileNotFoundError almost every + # barrier-synchronised round; the process-unique temp name makes every + # write an atomic last-writer-wins rename. + ctx = multiprocessing.get_context("spawn") + with tempfile.TemporaryDirectory() as tmpdir: + barrier = ctx.Barrier(2, timeout=60) + errors = ctx.Value("i", 0) + workers = [ + ctx.Process(target=_concurrent_writer, args=(tmpdir, barrier, 30, errors)) + for _ in range(2) + ] + for worker in workers: + worker.start() + for worker in workers: + worker.join(120) + + self.assertEqual(errors.value, 0) + entries = sorted(p.name for p in Path(tmpdir).iterdir()) + self.assertEqual(entries, ["ramtorch_prefetch_orders.json"]) + with (Path(tmpdir) / "ramtorch_prefetch_orders.json").open() as handle: + json.load(handle) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_segmented_checkpointing_model_support.py b/tests/test_segmented_checkpointing_model_support.py new file mode 100644 index 000000000..7180964bf --- /dev/null +++ b/tests/test_segmented_checkpointing_model_support.py @@ -0,0 +1,690 @@ +"""Model-level segmented checkpointing capability coverage.""" + +import unittest + + +def assert_checkpointing_controls( + test_case, + model_cls, + *, + backend=False, + interval=False, + stride=False, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, +): + if backend: + test_case.assertTrue(hasattr(model_cls, "set_gradient_checkpointing_backend")) + if interval: + test_case.assertTrue(hasattr(model_cls, "set_gradient_checkpointing_interval")) + if stride: + test_case.assertTrue(hasattr(model_cls, "set_gradient_checkpointing_segment_stride")) + if checkpoint_attention_offload: + test_case.assertTrue(hasattr(model_cls, "set_gradient_checkpointing_offload_attention")) + if ffn: + test_case.assertTrue(getattr(model_cls, "_supports_ffn_gradient_checkpointing", False)) + if attention_offload: + test_case.assertTrue(getattr(model_cls, "_supports_attention_activation_offload", False)) + + +class AceStepSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ace_step.transformer import ACEStepTransformer2DModel + + assert_checkpointing_controls( + self, + ACEStepTransformer2DModel, + backend=True, + interval=True, + stride=True, + ) + + +class AuraFlowSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.auraflow.transformer import AuraFlowTransformer2DModel + + assert_checkpointing_controls( + self, + AuraFlowTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class BooguImageSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.boogu_image.transformer import BooguImageTransformer2DModel + + assert_checkpointing_controls( + self, + BooguImageTransformer2DModel, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class ChromaSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.chroma.transformer import ChromaTransformer2DModel + + assert_checkpointing_controls( + self, + ChromaTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + def test_ffn_checkpointing_uses_non_reentrant_checkpoint(self): + import torch + + from simpletuner.helpers.models.chroma.transformer import ChromaTransformerBlock + + checkpoint_kwargs = {} + + def checkpoint_fn(function, *args, **kwargs): + checkpoint_kwargs.update(kwargs) + return function(*args) + + block = ChromaTransformerBlock(dim=16, num_attention_heads=2, attention_head_dim=8).train() + encoder_hidden_states, hidden_states = block( + hidden_states=torch.randn(2, 4, 16, requires_grad=True), + encoder_hidden_states=torch.randn(2, 3, 16, requires_grad=True), + image_temb=torch.randn(2, 6, 16), + text_temb=torch.randn(2, 6, 16), + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + ) + + self.assertEqual(checkpoint_kwargs, {"use_reentrant": False}) + self.assertEqual(encoder_hidden_states.shape, (2, 3, 16)) + self.assertEqual(hidden_states.shape, (2, 4, 16)) + + +class CosmosSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.cosmos.transformer import CosmosTransformer3DModel + + assert_checkpointing_controls( + self, + CosmosTransformer3DModel, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class Cosmos3SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.cosmos3.transformer import Cosmos3OmniTransformer + + assert_checkpointing_controls( + self, + Cosmos3OmniTransformer, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class ErnieSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ernie.transformer import ErnieImageTransformer2DModel + + assert_checkpointing_controls( + self, + ErnieImageTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class FluxSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.flux.transformer import FluxTransformer2DModel + + assert_checkpointing_controls( + self, + FluxTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + def test_checkpointing_forwards_attention_offload_to_block_wrappers(self): + from unittest.mock import patch + + import torch + import torch.nn as nn + + from simpletuner.helpers.models.flux.transformer import FluxTransformer2DModel + + class RecordingDoubleBlock(nn.Module): + def __init__(self): + super().__init__() + self.received_offload_attention = None + + def forward( + self, + hidden_states, + encoder_hidden_states, + temb, + context_temb=None, + image_rotary_emb=None, + attention_mask=None, + checkpoint_ffn=False, + checkpoint_fn=None, + offload_attention=False, + ): + self.received_offload_attention = offload_attention + return encoder_hidden_states + hidden_states.mean() * 0, hidden_states + temb.mean() * 0 + + class RecordingSingleBlock(nn.Module): + def __init__(self): + super().__init__() + self.received_offload_attention = None + + def forward( + self, + hidden_states, + temb, + image_rotary_emb=None, + attention_mask=None, + checkpoint_ffn=False, + checkpoint_fn=None, + offload_attention=False, + ): + self.received_offload_attention = offload_attention + return hidden_states + temb.mean() * 0 + + model = FluxTransformer2DModel( + patch_size=1, + in_channels=4, + num_layers=1, + num_single_layers=1, + attention_head_dim=6, + num_attention_heads=2, + joint_attention_dim=12, + pooled_projection_dim=12, + axes_dims_rope=(2, 2, 2), + ) + double_block = RecordingDoubleBlock() + single_block = RecordingSingleBlock() + model.transformer_blocks[0] = double_block + model.single_transformer_blocks[0] = single_block + model.train() + model.gradient_checkpointing = True + model.set_gradient_checkpointing_offload_attention(True) + + def fake_checkpoint(function, *args, **kwargs): + return function(*args) + + with patch("simpletuner.helpers.models.flux.transformer.simpletuner_checkpoint", side_effect=fake_checkpoint): + model( + hidden_states=torch.randn(1, 2, 4, requires_grad=True), + encoder_hidden_states=torch.randn(1, 3, 12), + pooled_projections=torch.randn(1, 12), + timestep=torch.tensor([1.0]), + img_ids=torch.zeros(2, 3), + txt_ids=torch.zeros(3, 3), + return_dict=True, + ) + + self.assertTrue(double_block.received_offload_attention) + self.assertTrue(single_block.received_offload_attention) + + +class FluxBlockCheckpointingScopeTests(unittest.TestCase): + def test_blocks_accept_ffn_checkpoint_and_attention_offload_scope(self): + import torch + + from simpletuner.helpers.models.flux.transformer import FluxSingleTransformerBlock, FluxTransformerBlock + + double_block = FluxTransformerBlock(dim=16, num_attention_heads=2, attention_head_dim=8).train() + hidden = torch.randn(2, 4, 16, requires_grad=True) + encoder_hidden = torch.randn(2, 3, 16, requires_grad=True) + temb = torch.randn(2, 16) + + expected_encoder, expected_hidden = double_block(hidden, encoder_hidden, temb) + actual_encoder, actual_hidden = double_block( + hidden, + encoder_hidden, + temb, + checkpoint_ffn=True, + checkpoint_fn=torch.utils.checkpoint.checkpoint, + offload_attention=True, + ) + self.assertTrue(torch.allclose(expected_encoder, actual_encoder, atol=1e-6)) + self.assertTrue(torch.allclose(expected_hidden, actual_hidden, atol=1e-6)) + + single_block = FluxSingleTransformerBlock(dim=16, num_attention_heads=2, attention_head_dim=8).train() + hidden = torch.randn(2, 7, 16, requires_grad=True) + temb = torch.randn(2, 16) + + expected_hidden = single_block(hidden, temb) + actual_hidden = single_block( + hidden, temb, checkpoint_ffn=True, checkpoint_fn=torch.utils.checkpoint.checkpoint, offload_attention=True + ) + self.assertTrue(torch.allclose(expected_hidden, actual_hidden, atol=1e-6)) + + +class Flux2SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.flux2.transformer import Flux2Transformer2DModel + + assert_checkpointing_controls( + self, + Flux2Transformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class HiDreamSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.hidream.transformer import HiDreamImageTransformer2DModel + + assert_checkpointing_controls( + self, + HiDreamImageTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class HunyuanVideoSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.hunyuanvideo.transformer import HunyuanVideo15Transformer3DModel + + assert_checkpointing_controls( + self, + HunyuanVideo15Transformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class IdeogramSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ideogram.transformer import Ideogram4Transformer + + assert_checkpointing_controls( + self, + Ideogram4Transformer, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class Kandinsky5SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.kandinsky5_video.transformer_kandinsky5 import Kandinsky5Transformer3DModel + + assert_checkpointing_controls( + self, + Kandinsky5Transformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class KolorsControlNetCheckpointingCompatibilityTests(unittest.TestCase): + def test_gradient_checkpointing_signature_accepts_diffusers_kwargs(self): + import inspect + + from simpletuner.helpers.models.kolors.controlnet import ControlNetModel + + parameters = inspect.signature(ControlNetModel._set_gradient_checkpointing).parameters + self.assertIn("enable", parameters) + self.assertIn("gradient_checkpointing_func", parameters) + + +class Krea2SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.krea2.transformer import Krea2Transformer2DModel + + assert_checkpointing_controls( + self, + Krea2Transformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + +class LongCatImageSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.longcat_image.transformer import LongCatImageTransformer2DModel + + assert_checkpointing_controls( + self, + LongCatImageTransformer2DModel, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class LongCatVideoSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.longcat_video.transformer import LongCatVideoTransformer3DModel + + assert_checkpointing_controls( + self, + LongCatVideoTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class LTXVideoSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ltxvideo.transformer import LTXVideoTransformer3DModel + + assert_checkpointing_controls( + self, + LTXVideoTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class LTXVideo2SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ltxvideo2.transformer import LTX2VideoTransformer3DModel + + assert_checkpointing_controls( + self, + LTX2VideoTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + +class Lumina2SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.lumina2.transformer import Lumina2Transformer2DModel + + assert_checkpointing_controls( + self, + Lumina2Transformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class MageFlowSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.mageflow.transformer import MageFlowTransformer2DModel + + assert_checkpointing_controls( + self, + MageFlowTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + +class PixArtSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.pixart.transformer import PixArtTransformer2DModel + + assert_checkpointing_controls( + self, + PixArtTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class QwenImageSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.qwen_image.transformer import QwenImageTransformer2DModel + + assert_checkpointing_controls( + self, + QwenImageTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class SanaSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.sana.transformer import SanaTransformer2DModel + + assert_checkpointing_controls( + self, + SanaTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class SanaVideoSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.sanavideo.transformer import SanaVideoTransformer3DModel + + assert_checkpointing_controls( + self, + SanaVideoTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class SD3SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.sd3.transformer import SD3Transformer2DModel + + assert_checkpointing_controls( + self, + SD3Transformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=False, + attention_offload=True, + ) + + +class StableCascadeSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.stable_cascade.unet import StableCascadeUNet + + assert_checkpointing_controls( + self, + StableCascadeUNet, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class WanSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.wan.transformer import WanTransformer3DModel + + assert_checkpointing_controls( + self, + WanTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + def test_block_accepts_ffn_checkpoint_and_attention_offload_scope(self): + import torch + + from simpletuner.helpers.models.wan.transformer import WanTransformerBlock + + torch.manual_seed(0) + block = WanTransformerBlock(dim=16, ffn_dim=64, num_heads=2).train() + hidden_states = torch.randn(2, 5, 16, requires_grad=True) + encoder_hidden_states = torch.randn(2, 7, 16, requires_grad=True) + temb = torch.randn(2, 6, 16) + rotary_emb = torch.randn(1, 1, hidden_states.shape[1], 4, dtype=torch.complex64) + + expected_hidden_states = block(hidden_states, encoder_hidden_states, temb, rotary_emb) + + checkpoint_use_reentrant_values = [] + + def checkpoint_fn(function, *args, **kwargs): + checkpoint_use_reentrant_values.append(kwargs.get("use_reentrant")) + return torch.utils.checkpoint.checkpoint(function, *args, **kwargs) + + actual_hidden_states = block( + hidden_states, + encoder_hidden_states, + temb, + rotary_emb, + checkpoint_ffn=True, + checkpoint_fn=checkpoint_fn, + offload_attention=True, + ) + + self.assertEqual(checkpoint_use_reentrant_values, [False]) + self.assertTrue(torch.allclose(expected_hidden_states, actual_hidden_states, atol=1e-6)) + + +class WanS2VSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.wan_s2v.transformer import WanS2VTransformer3DModel + + assert_checkpointing_controls( + self, + WanS2VTransformer3DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) + + +class ZImageSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.z_image.transformer import ZImageTransformer2DModel + + assert_checkpointing_controls( + self, + ZImageTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=True, + ffn=True, + attention_offload=True, + ) + + +class ZLabI1SegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.zlab_i1.transformer import ZlabI1Transformer2DModel + + assert_checkpointing_controls( + self, + ZlabI1Transformer2DModel, + backend=False, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + ) diff --git a/tests/test_stable_cascade_modules.py b/tests/test_stable_cascade_modules.py index 6161db259..448f044af 100644 --- a/tests/test_stable_cascade_modules.py +++ b/tests/test_stable_cascade_modules.py @@ -173,6 +173,38 @@ def test_unet_flowmap_from_config_restores_delta_blocks(self): self.assertEqual(clone.flowmap_deltatime_type, "r") self.assertTrue(has_delta_mapper) + def test_unet_gradient_checkpointing_interval_and_stride(self): + model = self._tiny_attention_unet() + sample, timestep_ratio, clip_text_pooled = self._tiny_unet_inputs() + + def checkpointed_block_names(interval=None, stride=None): + calls = [] + model.gradient_checkpointing = True + model.set_gradient_checkpointing_interval(interval) + model.set_gradient_checkpointing_segment_stride(stride) + + def fake_checkpoint(block, *args): + calls.append(block.__class__.__name__) + return block(*args) + + model._gradient_checkpointing_func = fake_checkpoint + model( + sample=sample, + timestep_ratio=timestep_ratio, + clip_text_pooled=clip_text_pooled, + ) + return calls + + self.assertEqual(len(checkpointed_block_names()), 6) + self.assertEqual( + checkpointed_block_names(interval=2), + ["SDCascadeResBlock", "SDCascadeAttnBlock", "SDCascadeTimestepBlock"], + ) + self.assertEqual( + checkpointed_block_names(interval=2, stride=4), + ["SDCascadeResBlock", "SDCascadeTimestepBlock", "SDCascadeTimestepBlock", "SDCascadeAttnBlock"], + ) + def _tiny_unet(self): return StableCascadeUNet( in_channels=4, @@ -196,6 +228,29 @@ def _tiny_unet(self): dropout=0.0, ) + def _tiny_attention_unet(self): + return StableCascadeUNet( + in_channels=4, + out_channels=4, + timestep_ratio_embedding_dim=8, + patch_size=1, + conditioning_dim=8, + block_out_channels=(8,), + num_attention_heads=(1,), + down_num_layers_per_block=(1,), + up_num_layers_per_block=(1,), + down_blocks_repeat_mappers=(1,), + up_blocks_repeat_mappers=(1,), + block_types_per_layer=(("SDCascadeResBlock", "SDCascadeTimestepBlock", "SDCascadeAttnBlock"),), + clip_text_pooled_in_channels=8, + clip_text_in_channels=None, + clip_image_in_channels=None, + clip_seq=1, + effnet_in_channels=None, + pixel_mapper_in_channels=None, + dropout=0.0, + ) + def _tiny_unet_inputs(self): sample = torch.randn(2, 4, 4, 4) timestep_ratio = torch.tensor([0.2, 0.4]) diff --git a/tests/test_trainer.py b/tests/test_trainer.py index 597c73e6c..b9aea9736 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -16,6 +16,7 @@ import torch +from simpletuner.helpers.publishing.huggingface import HubManager from simpletuner.helpers.publishing.providers.s3 import S3PublishingProvider from simpletuner.helpers.training.state_tracker import StateTracker from simpletuner.helpers.utils.checkpoint_manager import CheckpointManager @@ -1837,6 +1838,54 @@ def test_epoch_rollover(self, mock_state_tracker, mock_logger, mock_parse_args, self.assertEqual(trainer.state["current_epoch"], 2) self.assertEqual(trainer.extra_lr_scheduler_kwargs["epoch"], 2) + def test_epoch_rollover_reuses_initial_bucket_padding_policy(self): + cases = ( + ("fixed_steps", False, False, False), + ("epoch_driven", True, False, True), + ("oversubscribed_fixed_steps", False, True, True), + ) + + for name, overrode_max_train_steps, allow_oversubscription, expected in cases: + with self.subTest(name=name): + trainer = object.__new__(Trainer) + trainer.state = {"first_epoch": 1, "current_epoch": 1} + trainer.accelerator = MagicMock(is_main_process=True) + trainer.extra_lr_scheduler_kwargs = {} + trainer.get_steps_per_epoch_for_epoch = MagicMock(return_value=100) + trainer.config = SimpleNamespace( + num_train_epochs=5, + aspect_bucket_disable_rebuild=False, + lr_scheduler="constant", + num_update_steps_per_epoch=100, + gradient_accumulation_steps=3, + overrode_max_train_steps=overrode_max_train_steps, + allow_dataset_oversubscription=allow_oversubscription, + ) + metadata_backend = MagicMock(read_only=True) + backends = {"train": {"metadata_backend": metadata_backend}} + backend_config = { + "crop": True, + "crop_aspect": "random", + } + + with ( + patch("simpletuner.helpers.training.trainer.StateTracker.set_epoch"), + patch( + "simpletuner.helpers.training.trainer.StateTracker.get_data_backends", + return_value=backends, + ), + patch( + "simpletuner.helpers.training.trainer.StateTracker.get_data_backend_config", + return_value=backend_config, + ), + ): + trainer._epoch_rollover(2) + + metadata_backend.split_buckets_between_processes.assert_called_once_with( + gradient_accumulation_steps=3, + apply_padding=expected, + ) + @patch( "simpletuner.helpers.training.trainer.Trainer.parse_arguments", return_value=Mock(), @@ -2164,6 +2213,338 @@ def test_init_resume_checkpoint_deletes_unguarded_latest_when_guarded_checkpoint self.assertFalse(Path(tmpdir, "checkpoint-200").exists()) trainer.accelerator.load_state.assert_called_once_with(os.path.join(tmpdir, "checkpoint-100")) + def _build_standard_checkpoint_test_trainer(self, output_dir, limit, use_checkpoint_manager): + trainer = object.__new__(Trainer) + trainer.job_id = None + trainer.hub_manager = None + trainer.webhook_handler = None + trainer.validation = None + trainer._prepare_training_progress_payload = Mock(return_value=({}, {})) + trainer._emit_event = Mock() + trainer._run_post_upload_script = Mock() + trainer._drain_hub_upload_futures = Mock() + trainer.checkpoint_manager = CheckpointManager(output_dir) if use_checkpoint_manager else None + trainer.state = {"global_step": 110, "current_epoch": 1} + trainer.config = SimpleNamespace( + output_dir=output_dir, + use_deepspeed_optimizer=False, + fsdp_enable=False, + checkpoints_total_limit=limit, + num_train_epochs=1, + ) + trainer.accelerator = SimpleNamespace(is_main_process=True) + return trainer + + def test_checkpoint_state_cleanup_temp_removes_standard_and_rolling_temp(self): + for use_checkpoint_manager in (True, False): + with self.subTest(use_checkpoint_manager=use_checkpoint_manager): + with tempfile.TemporaryDirectory() as tmpdir: + for checkpoint in ( + "checkpoint-10", + "checkpoint-20-rolling", + "checkpoint-30-tmp", + "checkpoint-40-rolling-tmp", + ): + Path(tmpdir, checkpoint).mkdir() + + trainer = object.__new__(Trainer) + trainer.checkpoint_manager = CheckpointManager(tmpdir) if use_checkpoint_manager else None + trainer.checkpoint_state_cleanup_temp(tmpdir) + + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-10", "checkpoint-20-rolling"]) + + def test_standard_checkpoint_unlimited_cleans_temp_without_rotation(self): + for use_checkpoint_manager in (True, False): + with self.subTest(use_checkpoint_manager=use_checkpoint_manager): + with tempfile.TemporaryDirectory() as tmpdir: + existing_checkpoint = Path(tmpdir, "checkpoint-100") + incoming_checkpoint = Path(tmpdir, "checkpoint-110") + temp_checkpoint = Path(tmpdir, "checkpoint-90-tmp") + existing_checkpoint.mkdir() + temp_checkpoint.mkdir() + + trainer = self._build_standard_checkpoint_test_trainer( + tmpdir, + limit=0, + use_checkpoint_manager=use_checkpoint_manager, + ) + trainer.checkpoint_state_cleanup = Mock(wraps=trainer.checkpoint_state_cleanup) + + def save_checkpoint(output_dir): + self.assertEqual(tmpdir, output_dir) + self.assertFalse(temp_checkpoint.exists()) + incoming_checkpoint.mkdir() + return str(incoming_checkpoint) + + trainer.checkpoint_state_save = Mock(side_effect=save_checkpoint) + save_path = trainer._run_standard_checkpoint( + webhook_message=None, + parent_loss=None, + epoch=0, + upload_to_hub=False, + ) + + self.assertEqual(str(incoming_checkpoint), save_path) + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-100", "checkpoint-110"]) + trainer.checkpoint_state_cleanup.assert_not_called() + trainer._drain_hub_upload_futures.assert_not_called() + + def test_standard_checkpoint_waits_only_when_removal_is_needed(self): + for use_checkpoint_manager in (True, False): + with self.subTest(use_checkpoint_manager=use_checkpoint_manager): + with tempfile.TemporaryDirectory() as tmpdir: + existing_checkpoint = Path(tmpdir, "checkpoint-100") + incoming_checkpoint = Path(tmpdir, "checkpoint-110") + next_checkpoint = Path(tmpdir, "checkpoint-120") + existing_checkpoint.mkdir() + + trainer = self._build_standard_checkpoint_test_trainer( + tmpdir, + limit=2, + use_checkpoint_manager=use_checkpoint_manager, + ) + + def save_checkpoint(output_dir): + self.assertEqual(tmpdir, output_dir) + incoming_checkpoint.mkdir() + return str(incoming_checkpoint) + + trainer.checkpoint_state_save = Mock(side_effect=save_checkpoint) + save_path = trainer._run_standard_checkpoint( + webhook_message=None, + parent_loss=None, + epoch=0, + upload_to_hub=False, + ) + + self.assertEqual(str(incoming_checkpoint), save_path) + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-100", "checkpoint-110"]) + trainer._drain_hub_upload_futures.assert_not_called() + + def save_next_checkpoint(output_dir): + self.assertEqual(tmpdir, output_dir) + next_checkpoint.mkdir() + return str(next_checkpoint) + + trainer.state["global_step"] = 120 + trainer.checkpoint_state_save = Mock(side_effect=save_next_checkpoint) + save_path = trainer._run_standard_checkpoint( + webhook_message=None, + parent_loss=None, + epoch=0, + upload_to_hub=False, + ) + + self.assertEqual(str(next_checkpoint), save_path) + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-110", "checkpoint-120"]) + trainer._drain_hub_upload_futures.assert_called_once_with(wait=True) + + def test_standard_checkpoint_rotation_preserves_recovery_and_new_save(self): + for use_checkpoint_manager in (True, False): + with self.subTest(use_checkpoint_manager=use_checkpoint_manager): + with tempfile.TemporaryDirectory() as tmpdir: + resume_checkpoint = Path(tmpdir) / "checkpoint-100" + existing_checkpoint = Path(tmpdir) / "checkpoint-200" + incoming_checkpoint = Path(tmpdir) / "checkpoint-110" + temp_checkpoint = Path(tmpdir) / "checkpoint-90-tmp" + resume_checkpoint.mkdir() + existing_checkpoint.mkdir() + temp_checkpoint.mkdir() + + trainer = object.__new__(Trainer) + trainer.job_id = None + trainer.hub_manager = Mock() + trainer.hub_manager.upload_latest_checkpoint.return_value = ( + "remote/checkpoint-110", + str(incoming_checkpoint), + "repo-url", + ) + trainer.webhook_handler = None + trainer.validation = None + trainer._prepare_training_progress_payload = Mock(return_value=({}, {})) + trainer._emit_event = Mock() + trainer._run_post_upload_script = Mock() + trainer._schedule_hub_upload = Mock(side_effect=lambda _description, upload: upload()) + + def drain_uploads_before_removal(*, wait): + self.assertTrue(wait) + self.assertTrue(resume_checkpoint.exists()) + self.assertTrue(existing_checkpoint.exists()) + self.assertTrue(incoming_checkpoint.exists()) + + trainer._drain_hub_upload_futures = Mock(side_effect=drain_uploads_before_removal) + trainer.checkpoint_manager = CheckpointManager(tmpdir) if use_checkpoint_manager else None + trainer.state = {"global_step": 110, "current_epoch": 1} + trainer.config = SimpleNamespace( + output_dir=tmpdir, + use_deepspeed_optimizer=False, + fsdp_enable=False, + checkpoints_total_limit=1, + num_train_epochs=1, + resume_from_checkpoint="checkpoint-100", + ) + trainer.accelerator = SimpleNamespace(is_main_process=True) + + def fail_save(output_dir): + self.assertEqual(tmpdir, output_dir) + self.assertFalse(temp_checkpoint.exists()) + raise RuntimeError("save failed") + + trainer.checkpoint_state_save = Mock(side_effect=fail_save) + + with self.assertRaisesRegex(RuntimeError, "save failed"): + trainer._run_standard_checkpoint( + webhook_message=None, + parent_loss=None, + epoch=0, + upload_to_hub=False, + ) + self.assertTrue(resume_checkpoint.exists()) + self.assertTrue(existing_checkpoint.exists()) + + def save_checkpoint(output_dir): + self.assertEqual(tmpdir, output_dir) + self.assertTrue(resume_checkpoint.exists()) + self.assertTrue(existing_checkpoint.exists()) + incoming_checkpoint.mkdir() + return str(incoming_checkpoint) + + trainer.checkpoint_state_save = Mock(side_effect=save_checkpoint) + save_path = trainer._run_standard_checkpoint( + webhook_message=None, + parent_loss=None, + epoch=0, + upload_to_hub=True, + ) + + self.assertEqual(str(incoming_checkpoint), save_path) + self.assertTrue(Path(save_path).exists()) + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-110"]) + trainer.hub_manager.upload_latest_checkpoint.assert_called_once_with( + validation_images=None, + webhook_handler=None, + global_step=110, + epoch=1, + checkpoint_path=str(incoming_checkpoint), + ) + trainer._drain_hub_upload_futures.assert_called_once_with(wait=True) + + def test_hub_upload_uses_explicit_checkpoint_path(self): + with tempfile.TemporaryDirectory() as tmpdir: + Path(tmpdir, "checkpoint-200").mkdir() + incoming_checkpoint = Path(tmpdir, "checkpoint-110") + incoming_checkpoint.mkdir() + + hub_manager = object.__new__(HubManager) + hub_manager.config = SimpleNamespace(output_dir=tmpdir) + hub_manager.find_latest_checkpoint = Mock(return_value=Path(tmpdir, "checkpoint-200")) + hub_manager.upload_model = Mock(return_value="repo-url") + hub_manager._repo_url = Mock(return_value="remote/checkpoint-110") + + result = hub_manager.upload_latest_checkpoint( + validation_images=None, + webhook_handler=None, + global_step=110, + epoch=1, + checkpoint_path=str(incoming_checkpoint), + ) + + self.assertEqual(("remote/checkpoint-110", str(incoming_checkpoint), "repo-url"), result) + hub_manager.find_latest_checkpoint.assert_not_called() + hub_manager.upload_model.assert_called_once_with( + validation_images=None, + override_path=incoming_checkpoint, + webhook_handler=None, + global_step=110, + epoch=1, + ) + hub_manager._repo_url.assert_called_once_with("checkpoint-110") + + def test_rolling_checkpoint_rotation_preserves_recovery_and_new_save(self): + for use_checkpoint_manager in (True, False): + with self.subTest(use_checkpoint_manager=use_checkpoint_manager): + with tempfile.TemporaryDirectory() as tmpdir: + resume_checkpoint = Path(tmpdir, "checkpoint-100-rolling") + existing_checkpoint = Path(tmpdir, "checkpoint-200-rolling") + incoming_checkpoint = Path(tmpdir, "checkpoint-110-rolling") + temp_checkpoint = Path(tmpdir, "checkpoint-90-rolling-tmp") + resume_checkpoint.mkdir() + existing_checkpoint.mkdir() + temp_checkpoint.mkdir() + + trainer = object.__new__(Trainer) + trainer.checkpoint_manager = CheckpointManager(tmpdir) if use_checkpoint_manager else None + trainer.config = SimpleNamespace( + output_dir=tmpdir, + use_deepspeed_optimizer=False, + fsdp_enable=False, + checkpoints_rolling_total_limit=1, + ) + trainer.accelerator = SimpleNamespace(is_main_process=True) + + def fail_save(output_dir, suffix): + self.assertEqual(tmpdir, output_dir) + self.assertEqual("rolling", suffix) + self.assertFalse(temp_checkpoint.exists()) + raise RuntimeError("save failed") + + trainer.checkpoint_state_save = Mock(side_effect=fail_save) + with self.assertRaisesRegex(RuntimeError, "save failed"): + trainer._save_rolling_checkpoint() + self.assertTrue(resume_checkpoint.exists()) + self.assertTrue(existing_checkpoint.exists()) + + def save_checkpoint(output_dir, suffix): + self.assertEqual(tmpdir, output_dir) + self.assertEqual("rolling", suffix) + self.assertTrue(resume_checkpoint.exists()) + self.assertTrue(existing_checkpoint.exists()) + incoming_checkpoint.mkdir() + return str(incoming_checkpoint) + + trainer.checkpoint_state_save = Mock(side_effect=save_checkpoint) + save_path = trainer._save_rolling_checkpoint() + + self.assertEqual(str(incoming_checkpoint), save_path) + self.assertTrue(Path(save_path).exists()) + remaining = sorted(path.name for path in Path(tmpdir).glob("checkpoint-*")) + self.assertEqual(remaining, ["checkpoint-110-rolling"]) + + def test_rolling_checkpoint_non_main_save_ownership(self): + for use_deepspeed, fsdp_enable, should_save in ( + (False, False, False), + (True, False, True), + (False, True, True), + ): + with self.subTest(use_deepspeed=use_deepspeed, fsdp_enable=fsdp_enable): + trainer = object.__new__(Trainer) + trainer.config = SimpleNamespace( + output_dir="/tmp/output", + use_deepspeed_optimizer=use_deepspeed, + fsdp_enable=fsdp_enable, + checkpoints_rolling_total_limit=1, + ) + trainer.accelerator = SimpleNamespace(is_main_process=False) + trainer.checkpoint_state_cleanup_temp = Mock() + trainer.checkpoint_state_cleanup = Mock() + trainer.checkpoint_state_save = Mock(return_value="/tmp/output/checkpoint-110-rolling") + + save_path = trainer._save_rolling_checkpoint() + + if should_save: + self.assertEqual("/tmp/output/checkpoint-110-rolling", save_path) + trainer.checkpoint_state_save.assert_called_once_with("/tmp/output", "rolling") + else: + self.assertIsNone(save_path) + trainer.checkpoint_state_save.assert_not_called() + trainer.checkpoint_state_cleanup_temp.assert_not_called() + trainer.checkpoint_state_cleanup.assert_not_called() + @patch("simpletuner.helpers.training.trainer.AttentionBackendController.on_save_checkpoint") def test_checkpoint_state_save_writes_completion_guard(self, mock_attention_backend): with tempfile.TemporaryDirectory() as tmpdir: @@ -2216,6 +2597,68 @@ def save_state(path): self.assertNotIn(".guard", manifest["files"]) trainer.accelerator.wait_for_everyone.assert_not_called() + @patch("simpletuner.helpers.training.trainer.AttentionBackendController.on_save_checkpoint") + def test_checkpoint_state_save_synchronizes_only_all_rank_temp_writes(self, mock_attention_backend): + distributed_configs = ( + (DistributedType.DEEPSPEED, True, False, (True, False), ["wait", "save", "wait", "wait", "wait"]), + (DistributedType.FSDP, False, True, (True, False), ["wait", "save", "wait", "wait", "wait"]), + (DistributedType.MULTI_GPU, False, False, (True,), ["save"]), + ) + for distributed_type, use_deepspeed, fsdp_enable, process_roles, expected_events in distributed_configs: + for is_main_process in process_roles: + with self.subTest(distributed_type=distributed_type, is_main_process=is_main_process): + with tempfile.TemporaryDirectory() as tmpdir: + trainer = object.__new__(Trainer) + trainer.config = SimpleNamespace( + output_dir=tmpdir, + checkpointing_use_tempdir=True, + disk_low_threshold=None, + use_deepspeed_optimizer=use_deepspeed, + fsdp_enable=fsdp_enable, + ) + trainer.state = {"global_step": 100} + trainer.job_id = "test-job" + trainer.model = SimpleNamespace() + trainer.model_hooks = SimpleNamespace( + training_state_path="training_state.json", + get_modelspec_architecture=Mock(return_value="test/lora"), + ) + trainer.checkpoint_manager = CheckpointManager(tmpdir) + trainer.mark_optimizer_eval = Mock() + trainer.mark_optimizer_train = Mock() + trainer._emit_event = Mock() + trainer._run_post_checkpoint_script = Mock() + + events = [] + + def wait_for_everyone(): + events.append("wait") + + def save_state(path): + self.assertEqual(expected_events[: expected_events.index("save")], events) + events.append("save") + Path(path).mkdir(parents=True, exist_ok=True) + Path(path, "pytorch_lora_weights.safetensors").write_bytes(b"weights") + + trainer.accelerator = SimpleNamespace( + _models=[], + is_main_process=is_main_process, + distributed_type=distributed_type, + save_state=Mock(side_effect=save_state), + wait_for_everyone=Mock(side_effect=wait_for_everyone), + ) + + with patch( + "simpletuner.helpers.training.state_tracker.StateTracker.get_data_backends", + return_value={}, + ): + save_path = trainer.checkpoint_state_save(tmpdir) + + self.assertEqual(str(Path(tmpdir, "checkpoint-100")), save_path) + self.assertEqual(expected_events, events) + expected_written_path = Path(save_path) if is_main_process else Path(f"{save_path}-tmp") + self.assertTrue(expected_written_path.exists()) + @patch("simpletuner.helpers.training.trainer.logger") def test_init_resume_checkpoint_prodigy_without_split_groups(self, mock_logger): """Test that prodigy optimizer works with optimizers that don't have split_groups attribute""" @@ -2407,6 +2850,7 @@ def test_epoch_checkpoint_persists_next_epoch_without_state_mutation(self): trainer._emit_event = Mock() trainer._run_post_upload_script = Mock() trainer.checkpoint_state_cleanup = Mock() + trainer.checkpoint_manager = None trainer.state = {"global_step": 10, "global_resume_step": 0, "current_epoch": 2} trainer.config = SimpleNamespace( output_dir=str(tmpdir), diff --git a/tests/test_transformers/test_auraflow_transformer.py b/tests/test_transformers/test_auraflow_transformer.py index d1ee11df0..9da9dcc48 100644 --- a/tests/test_transformers/test_auraflow_transformer.py +++ b/tests/test_transformers/test_auraflow_transformer.py @@ -33,6 +33,8 @@ TypoTestUtils, ) +from simpletuner.helpers.models.auraflow.controlnet import AuraFlowControlNetModel + # Import the target classes from simpletuner.helpers.models.auraflow.transformer import ( AuraFlowFeedForward, @@ -912,6 +914,89 @@ def test_transformer_gradient_checkpointing_methods(self): model.set_gradient_checkpointing_interval(interval) self.assertEqual(model.gradient_checkpointing_interval, interval) + def test_gradient_checkpointing_handles_joint_and_single_blocks_backward(self): + model = AuraFlowTransformer2DModel( + sample_size=8, + patch_size=2, + in_channels=4, + out_channels=4, + num_mmdit_layers=1, + num_single_dit_layers=1, + attention_head_dim=8, + num_attention_heads=2, + joint_attention_dim=16, + caption_projection_dim=16, + pos_embed_max_size=16, + ) + model.train() + model.gradient_checkpointing = True + model.set_gradient_checkpointing_interval(None) + + hidden_states = torch.randn(1, 4, 8, 8, requires_grad=True) + encoder_hidden_states = torch.randn(1, 3, 16) + output = model( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + timestep=torch.tensor([1.0]), + return_dict=False, + )[0] + + output.float().mean().backward() + self.assertIsNotNone(hidden_states.grad) + + def test_controlnet_checkpointing_forwards_attention_kwargs(self): + class RecordingJointBlock(nn.Module): + def __init__(self): + super().__init__() + self.received_attention_kwargs = None + + def forward(self, hidden_states, encoder_hidden_states, temb, attention_kwargs=None): + self.received_attention_kwargs = attention_kwargs + return encoder_hidden_states + hidden_states.mean() * 0, hidden_states + temb.mean() * 0 + + class RecordingSingleBlock(nn.Module): + def __init__(self): + super().__init__() + self.received_attention_kwargs = None + + def forward(self, hidden_states, temb, attention_kwargs=None): + self.received_attention_kwargs = attention_kwargs + return hidden_states + temb.mean() * 0 + + model = AuraFlowControlNetModel( + sample_size=8, + patch_size=2, + in_channels=4, + out_channels=4, + num_mmdit_layers=1, + num_single_dit_layers=1, + attention_head_dim=8, + num_attention_heads=2, + joint_attention_dim=16, + caption_projection_dim=16, + pos_embed_max_size=16, + ) + joint_block = RecordingJointBlock() + single_block = RecordingSingleBlock() + model.joint_transformer_blocks[0] = joint_block + model.single_transformer_blocks[0] = single_block + model.train() + model.gradient_checkpointing = True + + hidden_states = torch.randn(1, 4, 8, 8, requires_grad=True) + attention_kwargs = {"processor_marker": "preserved"} + model( + hidden_states=hidden_states, + controlnet_cond=torch.zeros_like(hidden_states), + encoder_hidden_states=torch.randn(1, 3, 16), + timestep=torch.tensor([1.0]), + attention_kwargs=attention_kwargs, + return_dict=False, + ) + + self.assertEqual(joint_block.received_attention_kwargs, attention_kwargs) + self.assertEqual(single_block.received_attention_kwargs, attention_kwargs) + def test_transformer_tread_router_methods(self): """Test TREAD router configuration methods.""" with patch("diffusers.models.embeddings.Timesteps"), patch("diffusers.models.embeddings.TimestepEmbedding"): diff --git a/tests/test_transformers/test_pixart_transformer.py b/tests/test_transformers/test_pixart_transformer.py index 80cb7860a..3cab62f4a 100644 --- a/tests/test_transformers/test_pixart_transformer.py +++ b/tests/test_transformers/test_pixart_transformer.py @@ -901,6 +901,57 @@ def test_forward_rejects_wrong_tokenwise_timestep_length(self): return_dict=False, ) + def test_segmented_checkpointing_uses_sequential_state_helper(self): + from simpletuner.helpers.models.pixart.transformer import PixArtTransformer2DModel + + model = PixArtTransformer2DModel( + num_attention_heads=2, + attention_head_dim=8, + in_channels=4, + out_channels=8, + num_layers=4, + cross_attention_dim=16, + caption_channels=16, + sample_size=8, + patch_size=2, + num_embeds_ada_norm=1000, + use_additional_conditions=False, + ) + model.train() + model.gradient_checkpointing = True + model.gradient_checkpointing_interval = 2 + model.gradient_checkpointing_segment_stride = 4 + + def fake_checkpoint_sequential_state( + blocks, + segment_size, + state, + run_block, + checkpoint_fn, + checkpoint_kwargs=None, + segment_stride=None, + ): + return state + + with patch( + "simpletuner.helpers.models.pixart.transformer.checkpoint_sequential_state", + side_effect=fake_checkpoint_sequential_state, + ) as checkpoint_sequential: + output = model( + hidden_states=torch.randn(1, 4, 8, 8, requires_grad=True), + encoder_hidden_states=torch.randn(1, 3, 16), + timestep=torch.tensor([100]), + encoder_attention_mask=torch.ones(1, 3), + return_dict=False, + )[0] + + self.assertEqual(output.shape, (1, 8, 8, 8)) + checkpoint_sequential.assert_called_once() + args, kwargs = checkpoint_sequential.call_args + self.assertIs(args[0], model.transformer_blocks) + self.assertEqual(args[1], 2) + self.assertEqual(kwargs, {"segment_stride": 4}) + class TestPixArtTransformerIntegration(TransformerBaseTest): """Integration tests for PixArtTransformer2DModel with real components.""" diff --git a/tests/test_transformers/test_sanavideo_transformer.py b/tests/test_transformers/test_sanavideo_transformer.py index ac4d58620..0fc12efd3 100644 --- a/tests/test_transformers/test_sanavideo_transformer.py +++ b/tests/test_transformers/test_sanavideo_transformer.py @@ -1,4 +1,5 @@ import unittest +from unittest.mock import patch import torch @@ -158,6 +159,80 @@ def test_guidance_embeddings_reject_wrong_tokenwise_timestep_sign_length(self): timestep_sign=timestep_sign, ) + def test_segmented_checkpointing_uses_sequential_state_helper(self): + model = SanaVideoTransformer3DModel( + in_channels=4, + out_channels=4, + num_attention_heads=2, + attention_head_dim=8, + num_layers=4, + num_cross_attention_heads=2, + cross_attention_head_dim=8, + cross_attention_dim=16, + caption_channels=16, + mlp_ratio=2.0, + sample_size=2, + patch_size=(1, 1, 1), + qk_norm="rms_norm_across_heads", + rope_max_seq_len=32, + ) + model.train() + model.gradient_checkpointing = True + model.gradient_checkpointing_interval = 2 + model.gradient_checkpointing_segment_stride = 4 + + def fake_checkpoint_sequential_state( + blocks, + segment_size, + state, + run_block, + checkpoint_fn, + checkpoint_kwargs=None, + segment_stride=None, + ): + return state + + with patch( + "simpletuner.helpers.models.sanavideo.transformer.checkpoint_sequential_state", + side_effect=fake_checkpoint_sequential_state, + ) as checkpoint_sequential: + output = model( + hidden_states=torch.randn(1, 4, 2, 2, 2, requires_grad=True), + encoder_hidden_states=torch.randn(1, 5, 16), + timestep=torch.randint(0, 1000, (1, 8)), + ) + + output_tensor = output.sample if hasattr(output, "sample") else output + self.assertEqual(output_tensor.shape, (1, 4, 2, 2, 2)) + checkpoint_sequential.assert_called_once() + args, kwargs = checkpoint_sequential.call_args + self.assertIs(args[0], model.transformer_blocks) + self.assertEqual(args[1], 2) + self.assertEqual(kwargs, {"segment_stride": 4}) + + def test_unsloth_backend_uses_offloaded_checkpoint_for_per_block_path(self): + self.model.train() + self.model.gradient_checkpointing = True + self.model.gradient_checkpointing_backend = "unsloth" + self.model.gradient_checkpointing_interval = None + + def fake_offloaded_checkpoint(function, *args, **kwargs): + return function(*args, **kwargs) + + with patch( + "simpletuner.helpers.training.offloaded_gradient_checkpointer.offloaded_checkpoint", + side_effect=fake_offloaded_checkpoint, + ) as offloaded_checkpoint: + output = self.model( + hidden_states=torch.randn(1, 4, 2, 2, 2, requires_grad=True), + encoder_hidden_states=torch.randn(1, 5, 16), + timestep=torch.randint(0, 1000, (1, 8)), + ) + + output_tensor = output.sample if hasattr(output, "sample") else output + self.assertEqual(output_tensor.shape, (1, 4, 2, 2, 2)) + offloaded_checkpoint.assert_called_once() + if __name__ == "__main__": unittest.main() diff --git a/tests/test_transformers/test_sd3_transformer.py b/tests/test_transformers/test_sd3_transformer.py index 9aeb07129..c9ea33b27 100644 --- a/tests/test_transformers/test_sd3_transformer.py +++ b/tests/test_transformers/test_sd3_transformer.py @@ -15,6 +15,7 @@ import os import sys import unittest +from contextlib import nullcontext from typing import Any, Dict, List, Tuple from unittest.mock import MagicMock, Mock, patch @@ -233,6 +234,93 @@ def test_gradient_checkpointing_interval(self): mock_model.set_gradient_checkpointing_interval(interval) mock_model.set_gradient_checkpointing_interval.assert_called_once_with(interval) + def test_segmented_checkpointing_uses_sequential_state_helper(self): + from simpletuner.helpers.models.sd3.transformer import SD3Transformer2DModel + + model = SD3Transformer2DModel( + sample_size=8, + patch_size=2, + in_channels=4, + num_layers=4, + attention_head_dim=8, + num_attention_heads=2, + joint_attention_dim=32, + caption_projection_dim=16, + pooled_projection_dim=16, + out_channels=4, + pos_embed_max_size=8, + ) + model.train() + model.gradient_checkpointing = True + model.gradient_checkpointing_interval = 2 + model.gradient_checkpointing_segment_stride = 4 + + def fake_checkpoint_sequential_state( + blocks, + segment_size, + state, + run_block, + checkpoint_fn, + checkpoint_kwargs=None, + segment_stride=None, + ): + return state + + with patch( + "simpletuner.helpers.models.sd3.transformer.checkpoint_sequential_state", + side_effect=fake_checkpoint_sequential_state, + ) as checkpoint_sequential: + output = model( + hidden_states=torch.randn(1, 4, 8, 8, requires_grad=True), + encoder_hidden_states=torch.randn(1, 3, 32), + pooled_projections=torch.randn(1, 16), + timestep=torch.tensor([100]), + return_dict=False, + )[0] + + self.assertEqual(output.shape, (1, 4, 8, 8)) + checkpoint_sequential.assert_called_once() + args, kwargs = checkpoint_sequential.call_args + self.assertIs(args[0], model.transformer_blocks) + self.assertEqual(args[1], 2) + self.assertEqual(kwargs, {"segment_stride": 4}) + + def test_attention_offload_context_is_used(self): + from simpletuner.helpers.models.sd3.transformer import SD3Transformer2DModel + + model = SD3Transformer2DModel( + sample_size=8, + patch_size=2, + in_channels=4, + num_layers=1, + attention_head_dim=8, + num_attention_heads=2, + joint_attention_dim=32, + caption_projection_dim=16, + pooled_projection_dim=16, + out_channels=4, + pos_embed_max_size=8, + ) + model.train() + model.set_gradient_checkpointing_offload_attention(True) + + with patch( + "simpletuner.helpers.models.sd3.transformer.activation_offload_context", + side_effect=lambda *args, **kwargs: nullcontext(), + ) as offload_context: + output = model( + hidden_states=torch.randn(1, 4, 8, 8, requires_grad=True), + encoder_hidden_states=torch.randn(1, 3, 32), + pooled_projections=torch.randn(1, 16), + timestep=torch.tensor([100]), + return_dict=False, + )[0] + + self.assertEqual(output.shape, (1, 4, 8, 8)) + self.assertTrue(model.gradient_checkpointing_offload_attention) + offload_context.assert_called() + self.assertTrue(any(call.args and call.args[0] is True for call in offload_context.call_args_list)) + def test_tread_router_integration(self): """Test TREAD router setting and integration.""" with patch("simpletuner.helpers.models.sd3.transformer.SD3Transformer2DModel") as MockSD3: @@ -355,6 +443,8 @@ def test_typo_prevention_method_names(self): required_methods = [ "forward", "set_gradient_checkpointing_interval", + "set_gradient_checkpointing_segment_stride", + "set_gradient_checkpointing_offload_attention", "set_router", "enable_forward_chunking", "disable_forward_chunking", diff --git a/tests/test_vae.py b/tests/test_vae.py index 23cf5d98c..8dda79ad8 100644 --- a/tests/test_vae.py +++ b/tests/test_vae.py @@ -297,5 +297,108 @@ def encode(self, audio): self.assertEqual(output["latent_lengths"].shape[0], 1) +class TestMetadataFilterOnSplitShard(unittest.TestCase): + """A filtered sample must leave this rank's shard the same length as its DP peers'.""" + + IMAGES = [f"image-{index}.jpg" for index in range(5)] + + @classmethod + def _padded_shard_backend(cls, dp_rank, *, world_size=2, images=None): + from simpletuner.helpers.metadata.backends.base import MetadataBackend + from simpletuner.helpers.training.state_tracker import StateTracker + + backend = MagicMock(spec=MetadataBackend) + backend.id = "filtered-shard" + backend.batch_size = 1 + backend.repeats = 0 + backend.bucket_report = None + backend.dataset_type = DatasetType.IMAGE + backend.aspect_ratio_bucket_indices = {"1.0": list(images or cls.IMAGES)} + backend.read_only = False + backend.filtering_statistics = None + backend.accelerator = SimpleNamespace( + num_processes=world_size, + process_index=dp_rank, + is_main_process=dp_rank == 0, + ) + with ( + patch.dict("os.environ", {"SIMPLETUNER_SHUFFLE_BUCKETS": "0"}), + patch.object( + StateTracker, + "get_args", + return_value=SimpleNamespace(allow_dataset_oversubscription=True), + ), + patch.object(StateTracker, "get_data_backend_config", return_value={}), + ): + MetadataBackend.split_buckets_between_processes( + backend, + gradient_accumulation_steps=1, + apply_padding=True, + ) + return backend + + @staticmethod + def _cache(metadata_backend): + from simpletuner.helpers.caching.vae import VAECache + + cache = object.__new__(VAECache) + cache.id = "filtered-shard" + cache.metadata_backend = metadata_backend + return cache + + def test_padded_shard_keeps_its_length_when_a_duplicate_is_filtered(self): + backend = self._padded_shard_backend(1) + # Five images over two DP shards: this rank holds three slots, the last of which is a + # padding copy of the final global item. + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], ["image-3.jpg", "image-4.jpg", "image-4.jpg"]) + self.assertTrue(backend.read_only) + + self._cache(backend)._handle_metadata_filtered_sample(filepath="image-4.jpg", bucket="1.0", reason="problematic") + + shard = backend.aspect_ratio_bucket_indices["1.0"] + self.assertNotIn("image-4.jpg", shard) + self.assertEqual(len(shard), 3) + self.assertEqual(shard, ["image-3.jpg"] * 3) + + def test_padded_shard_keeps_its_length_when_a_unique_sample_is_filtered(self): + backend = self._padded_shard_backend(1) + self._cache(backend)._handle_metadata_filtered_sample(filepath="image-3.jpg", bucket="1.0", reason="nsfw") + + shard = backend.aspect_ratio_bucket_indices["1.0"] + self.assertNotIn("image-3.jpg", shard) + self.assertEqual(shard, ["image-4.jpg"] * 3) + + def test_a_shard_made_entirely_of_one_filtered_sample_cannot_be_refilled(self): + # Known limitation: with nothing left in the bucket there is no sample to repeat. The + # shortened cache is picked up by the next split. + backend = self._padded_shard_backend(4, world_size=5, images=[f"image-{index}.jpg" for index in range(9)]) + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], ["image-8.jpg", "image-8.jpg"]) + + self._cache(backend)._handle_metadata_filtered_sample(filepath="image-8.jpg", bucket="1.0", reason="problematic") + + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], []) + + def test_unsplit_backend_is_not_repadded(self): + # Control: before the split the backend holds the dataset, not a shard, so a filtered + # sample simply disappears. + backend = self._padded_shard_backend(1) + backend.aspect_ratio_bucket_indices = {"1.0": list(self.IMAGES)} + backend.read_only = False + + self._cache(backend)._handle_metadata_filtered_sample(filepath="image-4.jpg", bucket="1.0", reason="problematic") + + self.assertEqual(backend.aspect_ratio_bucket_indices["1.0"], self.IMAGES[:4]) + + def test_filter_action_is_still_queued_for_the_unsplit_cache(self): + backend = self._padded_shard_backend(1) + cache = self._cache(backend) + cache._handle_metadata_filtered_sample(filepath="image-4.jpg", bucket="1.0", reason="problematic") + + self.assertEqual( + [(action["filepath"], action["reason"]) for action in cache._deferred_metadata_filter_actions], + [("image-4.jpg", "problematic")], + ) + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_validation_audio_mock.py b/tests/test_validation_audio_mock.py index 5ee765722..7911478d4 100644 --- a/tests/test_validation_audio_mock.py +++ b/tests/test_validation_audio_mock.py @@ -1,12 +1,13 @@ import unittest from io import BytesIO +from types import SimpleNamespace from unittest.mock import MagicMock, patch import torch from simpletuner.helpers.models.common import AudioModelFoundation, ModelTypes, PipelineTypes, PredictionTypes from simpletuner.helpers.training import validation_audio -from simpletuner.helpers.training.validation import Validation, ValidationPrompt +from simpletuner.helpers.training.validation import Validation, ValidationPrompt, prepare_validation_prompt_list class MockAudioModel(AudioModelFoundation): @@ -74,6 +75,38 @@ def move_models(self, device): class TestAudioValidation(unittest.TestCase): + def test_raw_validation_prompts_do_not_require_text_embed_cache(self): + class NoTextCacheModel: + def uses_text_embeddings_cache(self): + return False + + def requires_conditioning_validation_inputs(self): + return False + + def should_precompute_validation_negative_prompt(self): + return False + + args = SimpleNamespace( + model_family="heartmula", + model_flavour="3b", + controlnet=False, + control=False, + validation_using_datasets=False, + validation_input=None, + validation_prompt_library=False, + user_prompt_library=None, + validation_prompt="driving pop, bright synths", + validation_negative_prompt="None", + validation_disable_unconditional=True, + data_backend_config="config/examples/heartmula-audio.json", + ) + + with patch("simpletuner.helpers.training.validation.StateTracker.get_args", return_value=args): + metadata = prepare_validation_prompt_list(args, embed_cache=None, model=NoTextCacheModel()) + + self.assertEqual([entry.prompt for entry in metadata["validation_prompts"]], [args.validation_prompt]) + self.assertEqual(metadata["validation_shortnames"], ["validation"]) + @patch("simpletuner.helpers.training.validation.StateTracker") @patch("simpletuner.helpers.training.validation.validation_audio.save_audio") @patch("simpletuner.helpers.training.validation.prepare_validation_prompt_list") diff --git a/tests/test_wan_model.py b/tests/test_wan_model.py index 9d6988b65..d21d0e0ab 100644 --- a/tests/test_wan_model.py +++ b/tests/test_wan_model.py @@ -2,8 +2,10 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +import torch + from simpletuner.helpers.models.common import PipelineTypes, VideoModelFoundation -from simpletuner.helpers.models.wan.model import Wan +from simpletuner.helpers.models.wan.model import Wan, add_first_frame_latent_conditioning class WanModelTests(unittest.TestCase): @@ -191,6 +193,52 @@ def test_unload_validation_models_clears_cached_peer_stages(self): super_unload.assert_called_once_with(model) self.assertEqual(model._wan_cached_stage_modules, {}) + def test_latent_i2v_conditioning_builds_36_channel_input(self): + latent_model_input = torch.zeros(1, 16, 3, 4, 5) + clean_latents = torch.arange(1 * 16 * 3 * 4 * 5, dtype=torch.float32).view(1, 16, 3, 4, 5) + vae = SimpleNamespace(config=SimpleNamespace(temperal_downsample=[1, 1])) + + conditioned = add_first_frame_latent_conditioning(latent_model_input, clean_latents, vae) + + self.assertEqual(tuple(conditioned.shape), (1, 36, 3, 4, 5)) + self.assertTrue(torch.equal(conditioned[:, :16], latent_model_input)) + self.assertTrue(torch.all(conditioned[:, 16:20, 0] == 1)) + self.assertTrue(torch.all(conditioned[:, 16:20, 1:] == 0)) + self.assertTrue(torch.equal(conditioned[:, 20:, :1], clean_latents[:, :, :1])) + self.assertTrue(torch.all(conditioned[:, 20:, 1:] == 0)) + + def test_latent_i2v_conditioning_uses_config_temporal_downsample(self): + latent_model_input = torch.zeros(1, 16, 3, 4, 5) + clean_latents = torch.zeros(1, 16, 3, 4, 5) + vae = SimpleNamespace(config=SimpleNamespace(temperal_downsample=[0, 1])) + + conditioned = add_first_frame_latent_conditioning(latent_model_input, clean_latents, vae) + + self.assertEqual(tuple(conditioned.shape), (1, 34, 3, 4, 5)) + self.assertTrue(torch.all(conditioned[:, 16:18, 0] == 1)) + self.assertTrue(torch.all(conditioned[:, 16:18, 1:] == 0)) + + def test_i2v_conditioning_uses_cached_latents_without_warning(self): + model = object.__new__(Wan) + model._is_i2v_like_flavour = MagicMock(return_value=False) + model._extract_conditioning_frames = MagicMock(return_value=(None, None)) + model.get_vae = MagicMock(return_value=SimpleNamespace(config=SimpleNamespace(temperal_downsample=[1, 1]))) + + hidden_states = torch.zeros(1, 16, 3, 4, 5) + clean_latents = torch.arange(1 * 16 * 3 * 4 * 5, dtype=torch.float32).view(1, 16, 3, 4, 5) + transformer_kwargs = {"hidden_states": hidden_states.clone()} + prepared_batch = {"is_i2v_data": True, "latents": clean_latents} + + with patch("simpletuner.helpers.models.wan.model.logger.warning") as warning: + Wan._apply_i2v_conditioning_to_kwargs(model, prepared_batch, transformer_kwargs) + + conditioned = transformer_kwargs["hidden_states"] + self.assertEqual(tuple(conditioned.shape), (1, 36, 3, 4, 5)) + self.assertTrue(torch.equal(conditioned[:, :16], hidden_states)) + self.assertTrue(torch.equal(conditioned[:, 20:, :1], clean_latents[:, :, :1])) + self.assertTrue(torch.all(conditioned[:, 20:, 1:] == 0)) + warning.assert_not_called() + if __name__ == "__main__": unittest.main() diff --git a/tests/test_zimage.py b/tests/test_zimage.py index f851763c9..21a9d7aa3 100644 --- a/tests/test_zimage.py +++ b/tests/test_zimage.py @@ -1,6 +1,7 @@ import inspect import json import unittest +from unittest.mock import patch import torch @@ -115,6 +116,35 @@ def test_gradient_checkpointing_flag_exists(self): model._set_gradient_checkpointing(enable=True) self.assertTrue(model.gradient_checkpointing) + def test_ffn_backend_disables_segmented_checkpointing(self): + model = ZImageTransformer2DModel( + all_patch_size=(2,), + all_f_patch_size=(1,), + in_channels=1, + dim=8, + n_layers=2, + n_refiner_layers=1, + n_heads=1, + n_kv_heads=1, + norm_eps=1e-5, + qk_norm=False, + cap_feat_dim=4, + rope_theta=1.0, + t_scale=1.0, + axes_dims=[2, 2, 4], + axes_lens=[64, 64, 64], + ) + model.train() + model._set_gradient_checkpointing(enable=True) + model.set_gradient_checkpointing_backend("torch-ffn") + model.set_gradient_checkpointing_interval(2) + + with patch("simpletuner.helpers.models.z_image.transformer.checkpoint_sequential_state") as segmented: + output = model([torch.zeros(1, 1, 8, 8, requires_grad=True)], torch.full((1,), 0.5), [torch.zeros(2, 4)])[0] + + segmented.assert_not_called() + self.assertEqual(output[0].shape, (1, 1, 8, 8)) + def test_context_parallel_only_marks_unified_transformer_blocks(self): model = ZImageTransformer2DModel( all_patch_size=(2,), diff --git a/tests/test_zlab_i1_model.py b/tests/test_zlab_i1_model.py index 89360a092..1130d09b1 100644 --- a/tests/test_zlab_i1_model.py +++ b/tests/test_zlab_i1_model.py @@ -11,6 +11,7 @@ from simpletuner.helpers.models.zlab_i1.model import ZLabI1 from simpletuner.helpers.models.zlab_i1.pipeline import ZlabI1Pipeline from simpletuner.helpers.models.zlab_i1.transformer import ZlabI1Transformer2DModel +from simpletuner.helpers.training.gradient_checkpointing_interval import checkpoint_sequential_state from simpletuner.helpers.training.layersync import LayerSyncRegularizer from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils import ramtorch as ramtorch_utils @@ -238,6 +239,31 @@ def test_transformer_supports_rectangular_aspect_bucket_latents(self): self.assertEqual(output.shape, latents.shape) self.assertTrue(torch.isfinite(output).all()) + def test_segmented_checkpointing_uses_sequential_state_helper(self): + self.transformer.train() + self.transformer.gradient_checkpointing = True + self.transformer.gradient_checkpointing_interval = 2 + self.transformer.gradient_checkpointing_segment_stride = 4 + self.transformer._gradient_checkpointing_func = lambda fn, *args, **kwargs: fn(*args) + + with patch( + "simpletuner.helpers.models.zlab_i1.transformer.checkpoint_sequential_state", + wraps=checkpoint_sequential_state, + ) as checkpoint_sequential: + output = self.transformer( + self.latents.detach().clone().requires_grad_(), + self.timesteps, + self.prompt_embeds, + self.attention_mask, + ) + + self.assertEqual(output.shape, self.latents.shape) + self.assertTrue(torch.isfinite(output).all()) + checkpoint_sequential.assert_called_once() + args, kwargs = checkpoint_sequential.call_args + self.assertEqual(args[1], 2) + self.assertEqual(kwargs, {"segment_stride": 4}) + @unittest.skipUnless(torch.cuda.is_available(), "CUDA not available") def test_musubi_streams_i1_blocks_on_cuda_forward(self): transformer = ZlabI1Transformer2DModel(