diff --git a/README.md b/README.md index b139feea..2163a6f4 100644 --- a/README.md +++ b/README.md @@ -250,6 +250,10 @@ Documentation sources are also available in [`docs/`](docs/). | ๐Ÿท๏ธ Model Name | ๐Ÿ“œ Reference Paper | ๐Ÿ“ฆ Source of Weights | |---------------|-------------------|---------------------| | Stable Diffusion 1.x (v1-1 to v1-5) | [High-Resolution Image Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) | `diffusers` | + | Stable Diffusion 2.x (2, 2-base, 2-1, 2-1-base, SD-Turbo) | [High-Resolution Image Synthesis with Latent Diffusion Models](https://arxiv.org/abs/2112.10752) | `diffusers` | + | Stable Diffusion XL (base + refiner 1.0, SDXL-Turbo) | [SDXL: Improving Latent Diffusion Models for High-Resolution Image Synthesis](https://arxiv.org/abs/2307.01952) | `diffusers` | + | Stable Diffusion 3 (medium) | [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206) | `diffusers` | + | Stable Diffusion 3.5 (large, large-turbo, medium) | [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206) | `diffusers` |
@@ -276,7 +280,7 @@ Documentation sources are also available in [`docs/`](docs/). ## ๐Ÿ“œ License -This project leverages [timm](https://github.com/huggingface/pytorch-image-models#licenses), [transformers](https://github.com/huggingface/transformers#license) and [diffusers](https://github.com/huggingface/diffusers#license) for converting pretrained weights from PyTorch to Keras. For licensing details, please refer to the respective repositories. Converted weights keep their upstream license (for example, the Stable Diffusion checkpoints are CreativeML OpenRAIL-M). +This project leverages [timm](https://github.com/huggingface/pytorch-image-models#licenses), [transformers](https://github.com/huggingface/transformers#license) and [diffusers](https://github.com/huggingface/diffusers#license) for converting pretrained weights from PyTorch to Keras. For licensing details, please refer to the respective repositories. Converted weights keep their upstream license (for example, the Stable Diffusion checkpoints are CreativeML OpenRAIL-M / OpenRAIL++-M, and SDXL-Turbo is non-commercial under the Stability AI Community License). - ๐Ÿ”– **zeromodels Code**: This repository is licensed under the [Apache 2.0 License](https://www.apache.org/licenses/LICENSE-2.0). diff --git a/assets/stable_diffusion_2_ramen.jpg b/assets/stable_diffusion_2_ramen.jpg new file mode 100644 index 00000000..9c2f8225 Binary files /dev/null and b/assets/stable_diffusion_2_ramen.jpg differ diff --git a/assets/stable_diffusion_2_sd_turbo_raccoon.jpg b/assets/stable_diffusion_2_sd_turbo_raccoon.jpg new file mode 100644 index 00000000..08f2086f Binary files /dev/null and b/assets/stable_diffusion_2_sd_turbo_raccoon.jpg differ diff --git a/assets/stable_diffusion_2_text_to_image_output.jpg b/assets/stable_diffusion_2_text_to_image_output.jpg new file mode 100644 index 00000000..eaa2e3e2 Binary files /dev/null and b/assets/stable_diffusion_2_text_to_image_output.jpg differ diff --git a/assets/stable_diffusion_3_5_treehouse.jpg b/assets/stable_diffusion_3_5_treehouse.jpg new file mode 100644 index 00000000..c20622b7 Binary files /dev/null and b/assets/stable_diffusion_3_5_treehouse.jpg differ diff --git a/assets/stable_diffusion_3_corgi.jpg b/assets/stable_diffusion_3_corgi.jpg new file mode 100644 index 00000000..862f1fbe Binary files /dev/null and b/assets/stable_diffusion_3_corgi.jpg differ diff --git a/assets/stable_diffusion_xl_lighthouse.jpg b/assets/stable_diffusion_xl_lighthouse.jpg new file mode 100644 index 00000000..8abe970a Binary files /dev/null and b/assets/stable_diffusion_xl_lighthouse.jpg differ diff --git a/assets/stable_diffusion_xl_refiner_lighthouse.jpg b/assets/stable_diffusion_xl_refiner_lighthouse.jpg new file mode 100644 index 00000000..7797b843 Binary files /dev/null and b/assets/stable_diffusion_xl_refiner_lighthouse.jpg differ diff --git a/assets/stable_diffusion_xl_text_to_image_output.jpg b/assets/stable_diffusion_xl_text_to_image_output.jpg new file mode 100644 index 00000000..521ba9b9 Binary files /dev/null and b/assets/stable_diffusion_xl_text_to_image_output.jpg differ diff --git a/assets/stable_diffusion_xl_turbo_raccoon.jpg b/assets/stable_diffusion_xl_turbo_raccoon.jpg new file mode 100644 index 00000000..02d5bb37 Binary files /dev/null and b/assets/stable_diffusion_xl_turbo_raccoon.jpg differ diff --git a/docs/main_classes.md b/docs/main_classes.md index 84fa48e1..64bad0b9 100644 --- a/docs/main_classes.md +++ b/docs/main_classes.md @@ -248,15 +248,27 @@ model.generate( guidance_scale=None, seed=None, latents=None, + image=None, + strength=None, + denoising_start=None, + denoising_end=None, + output_type="image", + **conditioning, ) ``` -The diffusion flavor, used by [Stable Diffusion](stable_diffusion.md). Where the LM +The diffusion flavor, used by [Stable Diffusion](stable_diffusion.md), +[Stable Diffusion 2](stable_diffusion_2.md), [Stable Diffusion XL](stable_diffusion_xl.md) +and [Stable Diffusion 3 / 3.5](stable_diffusion_3.md). Where the LM mixins decode tokens, this one runs a scheduler's denoising loop with classifier-free guidance over a latent and decodes it to `(batch, H, W, 3)` uint8 images. A model supplies five hooks (`encode_prompt`, `unconditional_ids`, `predict_noise`, `decode_latents`, `latent_shape`) and a `scheduler`; the mixin owns the guidance batching, -the initial latent, the loop and the postprocess. The denoiser call is compiled per +the initial latent, the loop and the postprocess. `encode_prompt` may return a nested +structure of tensors rather than one (SDXL's context, pooled embedding and size ids), +which the guidance batching and the compiled step carry through unchanged; a model whose +unconditional branch is not an encoded prompt overrides `encode_negative_prompt` (SDXL +zeroes it). The denoiser call is compiled per backend (`jax.jit`, `tf.function(jit_compile=True)`, eager on Torch) and cached on the instance, like the LM decode loop. @@ -270,17 +282,34 @@ instance, like the LM decode loop. guidance off. - **seed** (`int`, *optional*): seeds the initial latent, reproducible per backend. - **latents** (*optional*): an explicit initial latent, for results identical across - backends. + backends; with `strength` or `denoising_start`, the clean latent to start from. +- **image** / **strength** (*optional*): image-to-image (diffusers' `Img2ImgPipeline`): + the `(batch, H, W, 3)` uint8 or `[0, 1]` float image is VAE-encoded (the + `encode_latents` hook), noised to the `strength` point of the schedule and the + remaining steps are run; `strength` defaults to the repo's `generate_args` (0.8, the + SDXL refiner 0.3). +- **denoising_end** / **denoising_start** (*optional*): stop after, or resume from, a + fraction of the schedule with no noise added: the SDXL base + refiner ensemble + (`output_type="latent"` hands the base's latent over). +- **output_type** (*optional*): `"image"` (uint8 images) or `"latent"`. +- **conditioning** (*optional*): any further keyword argument is model-specific + conditioning handed to `encode_prompt` (SDXL's `original_size` / + `crops_coords_top_left` / `target_size`); a `negative_` twin applies to the + negative branch only and defaults to the positive value, the way `negative_input_ids` + pairs with `input_ids`. ### BaseScheduler The samplers a diffusion model steps with (`PNDMScheduler`, `DDIMScheduler`, -`EulerDiscreteScheduler`, `EulerAncestralDiscreteScheduler` in +`EulerDiscreteScheduler`, `EulerAncestralDiscreteScheduler`, +`FlowMatchEulerDiscreteScheduler` (rectified flow, SD 3) in `zeromodels.base.base_scheduler`) are weightless classes with the diffusers interface: `set_timesteps(n)`, `scale_model_input(sample, t)`, `step(noise_pred, t, sample)` and `init_noise_sigma`. `get_scheduler(config)` builds the one a diffusers scheduler config dict names, which is what a repo's `scheduler_config` goes through; `model.scheduler` can -be swapped between calls. +be swapped between calls. The Euler samplers take the diffusers `timestep_spacing` +(`linspace`, `leading`, `trailing`) and `interpolation_type` (`linear`, `log_linear`), and +the schedules (betas, `alphas_cumprod`, sigmas, timesteps) match diffusers to the bit. ## Preprocessing @@ -389,6 +418,33 @@ layer in the library so a single implementation choice applies everywhere. - **scale** (`float`): the `1/sqrt(head_dim)` factor, applied inside. - **attention_mask** (*optional*): additive mask broadcastable to `(B, heads, T_q, T_kv)`. - **soft_cap** (`float`, *optional*): logit soft-capping, used by Gemma 2. -- **attn_implementation** (`str`, *optional*): pick the kernel; pass it through `from_weights` to set it model-wide. +- **attn_implementation** (`str`, *optional*): `"sdpa"`, `"fused"` or `"flash"` (below); `None` uses the implementation active in the current context, else the layer's own default. + +**Implementations.** All three compute the same `softmax(QKแต€ ยท scale + mask) ยท V`; they differ in +memory, speed and portability: + +| | `"sdpa"` (library default) | `"fused"` | `"flash"` | +|---|---|---|---| +| What runs | hand-written `matmul`, float32 `softmax`, `matmul` | `keras.ops.dot_product_attention`, the backend picks the kernel | the same op with `flash_attention=True` | +| torch | the math on any device / dtype | `scaled_dot_product_attention`: flash or memory-efficient kernel on a CUDA GPU in fp16 / bf16, else its math kernel | the flash kernel, or an error | +| JAX | the math | the XLA reference implementation (same memory as the math) | cuDNN flash on a capable GPU, or an error | +| TensorFlow | the math | falls back to the math | falls back to the math | +| Logits memory | the full `(B, heads, T_q, T_kv)` matrix, materialized in fp16 and again in float32 for the softmax | tiled, never materialized when a fused kernel applies | tiled | +| Masks / soft-cap / dropout | all supported | additive masks yes; a soft-cap or attention dropout falls back to the math | no masks; soft-cap / dropout fall back | + +The math path is what every parity number in this library was measured with and what runs +everywhere identically. Its cost is quadratic memory: at 1024px the Stable Diffusion 3 joint +attention (4429 tokens) needs about 3.8 GB of float32 logits per block, which runs an 8 GB GPU +out of memory, while `"fused"` runs the same model at a 6.9 GB peak, 1.0 s/step on an RTX +4060 Laptop (the SDXL 1024px step drops from a 10.3 GB to a 7.6 GB peak). The fused kernels +accumulate in float32 and differ from the math only by rounding (about 3e-7 in float32). + +**Precedence.** `Model.from_weights(attn_implementation=...)` activates the choice for the +model's build and every forward / generation step (a `ContextVar`, restored on exit, so it +never leaks to another model). A layer may carry its own default for when nothing is +chosen: the SD 3 MMDiT's `StableDiffusion3JointAttention` defaults to `"fused"` because of the sequence +length above; every other layer defaults to `"sdpa"`. An explicit choice always wins over a +layer default. Outside `from_weights`, `zeromodels.base.base_attention.use_attn_implementation("fused")` +wraps a build or a forward the same way. See also [Utilities](utils.md) for the image, video, visualization, and label helpers. diff --git a/docs/models.md b/docs/models.md index a8e3b412..f5328f81 100644 --- a/docs/models.md +++ b/docs/models.md @@ -106,3 +106,7 @@ Vision-language encoders, generative VLMs, grounding across detection, OCR, poin - [SigLIP](siglip.md) - [SigLIP 2](siglip2.md) - [Stable Diffusion](stable_diffusion.md) +- [Stable Diffusion 2](stable_diffusion_2.md) +- [Stable Diffusion XL](stable_diffusion_xl.md) +- [Stable Diffusion 3](stable_diffusion_3.md) +- [Stable Diffusion 3.5](stable_diffusion_3_5.md) diff --git a/docs/stable_diffusion.md b/docs/stable_diffusion.md index 96b0bfe2..2993961d 100644 --- a/docs/stable_diffusion.md +++ b/docs/stable_diffusion.md @@ -28,10 +28,11 @@ Key facts of the port: Weights are layout-independent, so one hosted checkpoint serves both, and the converter is a `(O, I, H, W) -> (H, W, I, O)` kernel transpose and nothing else. - **Block-level layers**: the UNet and VAE are built from composite Keras layers - (`ResnetBlock2D`, `Transformer2DModel`, `CrossAttention`, ...). A functional graph keeps - every node's output alive until the forward ends, and at 512px the UNet's per-op - intermediates (4096x4096 attention maps above all) would need several GB; a layer's - internals are freed when its call returns. + (`StableDiffusionResnetBlock2D`, `StableDiffusionTransformer2DModel`, + `StableDiffusionCrossAttention`, ...). A functional graph keeps every node's output + alive until the forward ends, and at 512px the UNet's per-op intermediates (4096x4096 + attention maps above all) would need several GB; a layer's internals are freed when + its call returns. - **Schedulers match diffusers to the bit**: the training noise schedule is built in float32 the way torch does it, so `alphas_cumprod` and the Euler sigmas are identical and a 50-step run stays within float rounding of the reference. @@ -347,4 +348,4 @@ The five hosted checkpoints are the supported weights; any repo laid out like th loads with `from_weights("/")`. The `hf:` prefix raises for diffusion models: convert a diffusers-format checkpoint once with `zeromodels/models/stable_diffusion/convert_stable_diffusion_diffusers_to_keras.py` -(`build_from_diffusers(repo)`, `pip install zeromodels[conversion]`) and host the result. +(`transfer_stable_diffusion(repo)`, `pip install zeromodels[conversion]`) and host the result. diff --git a/docs/stable_diffusion_2.md b/docs/stable_diffusion_2.md new file mode 100644 index 00000000..f08fb024 --- /dev/null +++ b/docs/stable_diffusion_2.md @@ -0,0 +1,223 @@ +# Stable Diffusion 2 + +
+Weights: pretrained Keras weights live on Hugging Face under +zeromodels/<variant> +(each repo carries zm_config.json + model.weights.h5 + +tokenizer.json). Load with from_weights("zeromodels/<variant>"). +
+ +Stable Diffusion 2.x, ported to pure Keras 3: the [Stable Diffusion](stable_diffusion.md) +latent-diffusion architecture with the second-generation conditioning. It reuses the +SD 1.x family end to end (`StableDiffusion2Model` is `StableDiffusionModel`, +`StableDiffusion2TextToImage` is `StableDiffusionTextToImage`, each with the SD 2 +configuration), so everything on that page applies: one container per repo, `generate` +from `BaseDiffusion`, the schedulers, both data formats. What changes is the config: + +- **Text encoder**: OpenCLIP ViT-H/14's text tower, 1024 wide, 16 heads, `gelu`, shipped + truncated to its 23rd (penultimate) layer, which SD 2 conditions on. +- **UNet**: 1024-d cross-attention, one head count per level (`(5, 10, 20, 20)`, a 64-wide + head everywhere, where SD 1.x used 8 heads at every level) and a linear token + projection in the Transformer2D blocks instead of the 1x1 conv. +- **Tokenizer**: the same CLIP BPE, padded with `!` (id 0) the OpenCLIP way rather than + `<|endoftext|>`, so the empty prompt used for classifier-free guidance is + `[<|startoftext|>, <|endoftext|>, !, !, ...]`. +- **768px checkpoints** (`stable-diffusion-2`, `stable-diffusion-2-1`): built at a 96x96 + latent and trained with the **v-prediction** objective; their repos carry a + v-prediction DDIM scheduler, which `generate` picks up from `scheduler_config`. +- **SD-Turbo** (`sd-turbo`): SD 2.1 distilled with Adversarial Diffusion Distillation to + generate in 1 to 4 steps **without guidance**; its repo carries an Euler scheduler with + `trailing` timestep spacing and `generate_args` of 1 step, `guidance_scale=0.0`. + +The weights are converted once, offline, and hosted: on-the-fly `hf:` conversion is +deliberately **not supported** for diffusion models. The original `stabilityai/*` repos are +no longer on the Hub; the conversion sources are the `sd2-community` mirrors. + +Links: + +- Paper: [High-Resolution Image Synthesis with Latent Diffusion Models (arXiv:2112.10752)](https://arxiv.org/abs/2112.10752) +- Reference implementation: [diffusers `StableDiffusionPipeline`](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/text2img) +- License: [CreativeML Open RAIL++-M](https://huggingface.co/sd2-community/stable-diffusion-2-1/blob/main/LICENSE-MODEL) + +See also [stable_diffusion.md](stable_diffusion.md), [clip.md](clip.md). + +## Variants + +Preconverted, float32 weights are hosted under `zeromodels/`. Load with +`from_weights("zeromodels/")`. Each repo is one container: UNet (866M) + VAE +(84M) + OpenCLIP text encoder (340M), 1.29B parameters, 5.16 GB (4.81 GiB, one +`model.weights.h5`). The SD 2 checkpoints are released under the CreativeML Open RAIL++-M +license; SD-Turbo under the Stability AI Non-Commercial Research Community License. + +| Variant | Hub | Resolution | Objective | Training | +|---|---|---|---|---| +| `stable-diffusion-2-base` | [`zeromodels/stable-diffusion-2-base`](https://huggingface.co/zeromodels/stable-diffusion-2-base) | 512 | epsilon | from scratch: 550k steps at 256px on LAION-5B (aesthetics >= 4.5), then 850k steps at 512px | +| `stable-diffusion-2` | [`zeromodels/stable-diffusion-2`](https://huggingface.co/zeromodels/stable-diffusion-2) | 768 | v-prediction | 2-base + 150k steps at 768px | +| `stable-diffusion-2-1-base` | [`zeromodels/stable-diffusion-2-1-base`](https://huggingface.co/zeromodels/stable-diffusion-2-1-base) | 512 | epsilon | 2-base + 220k steps at 512px (punsafe 0.98) | +| `stable-diffusion-2-1` | [`zeromodels/stable-diffusion-2-1`](https://huggingface.co/zeromodels/stable-diffusion-2-1) | 768 | v-prediction | 2 + 55k steps (punsafe 0.1) + 155k steps (punsafe 0.98) at 768px | +| `sd-turbo` | [`zeromodels/sd-turbo`](https://huggingface.co/zeromodels/sd-turbo) | 512 | epsilon, 1 to 4 steps, no guidance | SD 2.1 distilled with Adversarial Diffusion Distillation (non-commercial) | + +Use the `-base` checkpoints for 512px images and the others for 768px; each repo's +`zm_config.json` builds the graph at its native size. + +## API + +Configs are typed: `StableDiffusion2Config` (composite, `model_type` +`"stable_diffusion_2"`) over `StableDiffusion2UNetConfig` (a `UNet2DConditionConfig` +with the SD 2 widths), `AutoencoderKLConfig` and `StableDiffusion2TextConfig` (a +`CLIPTextConfig` with the ViT-H/14 sizes), plus the scheduler config and the token ids +(`pad_token_id` 0). Flat constructor, `unet_` / `vae_` / `text_` prefixes, like SD 1.x. + +### `StableDiffusion2TextToImage` + +`StableDiffusionTextToImage` with the SD 2 configuration; `generate` is unchanged: + +```python +generate( + input_ids, + attention_mask=None, + negative_input_ids=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, +) +``` + +Returns `(batch, H, W, 3)` uint8 images at the checkpoint's resolution; `latents` is +`(batch, 64, 64, 4)` for the 512px checkpoints and `(batch, 96, 96, 4)` for the 768px +ones (`channels_first`: channels second). Defaults come from the repo's `generate_args` +(50 steps, guidance 7.5). See [Stable Diffusion](stable_diffusion.md#stablediffusiontexttoimage) +for the argument table. + +### `StableDiffusion2Model` + +The container, `StableDiffusionModel` with the SD 2 configuration: the same three +disconnected paths (`.unet`, `.vae`, `.text_encoder`) and the same inputs / outputs, with +`encoder_hidden_states` 1024 wide. It loads the same repo as the task class. + +The components are the SD 1.x classes, `UNet2DConditionModel` (built with +`num_attention_heads=(5, 10, 20, 20)`, `use_linear_projection=True`, +`cross_attention_dim=1024`) and `AutoencoderKL`; see their tables on the +[Stable Diffusion](stable_diffusion.md#api) page. + +## Preprocessing + +### `StableDiffusion2Tokenizer` + +`StableDiffusionTokenizer` with `!` as the pad token: CLIP BPE, `<|startoftext|>` / +`<|endoftext|>` framing, truncated and `!`-padded to 77 tokens. Returns +`{"input_ids", "attention_mask"}`. + +```python +StableDiffusion2Tokenizer( + hf_id=None, tokenizer_file=None, max_seq_len=77, pad_token="!" +) +``` + +## End-to-end example + +```python +import os + +os.environ["KERAS_BACKEND"] = "torch" # or "jax" / "tensorflow" + +from PIL import Image +from zeromodels.models.stable_diffusion_2 import ( + StableDiffusion2TextToImage, + StableDiffusion2Tokenizer, +) + +model = StableDiffusion2TextToImage.from_weights("zeromodels/stable-diffusion-2-1-base") +tokenizer = StableDiffusion2Tokenizer.from_weights( + "zeromodels/stable-diffusion-2-1-base" +) + +inputs = tokenizer( + "a steaming bowl of ramen on a wooden table, food photography, shallow depth of field" +) +images = model.generate(**inputs, num_inference_steps=50, guidance_scale=7.5, seed=2) + +Image.fromarray(images[0]).save("ramen.png") # (512, 512, 3) uint8 +``` + +Stable Diffusion 2.1-base: a steaming bowl of ramen on a wooden table, 512px + +### 768px, v-prediction + +The 768px checkpoints need no extra arguments; the v-prediction DDIM scheduler and the +96x96 latent come from the repo: + +```python +model = StableDiffusion2TextToImage.from_weights("zeromodels/stable-diffusion-2-1") +tokenizer = StableDiffusion2Tokenizer.from_weights("zeromodels/stable-diffusion-2-1") +images = model.generate( + **tokenizer("a lighthouse on a cliff at dusk, oil painting") +) # (1, 768, 768, 3) +``` + +### SD-Turbo + +One step, no guidance (the repo's defaults; up to 4 steps sharpen a little): + +```python +model = StableDiffusion2TextToImage.from_weights("zeromodels/sd-turbo") +tokenizer = StableDiffusion2Tokenizer.from_weights("zeromodels/sd-turbo") +images = model.generate( + **tokenizer( + "a cinematic shot of a baby raccoon wearing an intricate italian priest robe" + ) +) +``` + +SD-Turbo, one step: a baby raccoon in an italian priest robe, 512px + +The same prompt and latent through diffusers' `StableDiffusionPipeline` (fp32) match to +1 uint8 level: 99.9% of pixels identical at 1 step (PSNR 82 dB), 99.8% at 4 steps. + +Negative prompts, batching, explicit `latents` for cross-backend reproducibility, other +resolutions and the container-only use work exactly as on the +[Stable Diffusion](stable_diffusion.md#end-to-end-example) page. Image-to-image +(`generate(..., image=..., strength=...)`) works as described for +[`BaseDiffusion`](main_classes.md#basediffusion). + +### Verified against diffusers + +The ramen prompt above with the same initial latent, PNDM 50 steps and guidance 7.5 +through `StableDiffusion2TextToImage` (left) and diffusers' `StableDiffusionPipeline` +(right) on `stable-diffusion-2-1-base`, both fp32: + +zeromodels (left) and diffusers (right) generations of a bowl of ramen from stable-diffusion-2-1-base with the same latent + +``` +tokenizer input_ids identical +text_encoder max|d|=1.5e-05 +unet noise_pred (t=981) max|d|=7.2e-07 +vae decode max|d|=2.6e-06 +final image (uint8) max|d|=1 mean|d|=0.0095 PSNR=68.3 dB identical_pixels=97.22% +``` + +The 768px v-prediction checkpoint (`stable-diffusion-2-1`, DDIM 20 steps, same latent) +matches the same way: UNet 8.3e-6, VAE 7.9e-5, final image max 1 uint8 level, 99.3% of +pixels identical, PSNR 74.6 dB. + +## Data Format + +Both `channels_last` and `channels_first` are supported, as for +[Stable Diffusion](stable_diffusion.md#data-format); `generate` always returns +`(batch, H, W, 3)` uint8. + +## Memory and speed + +The fp32 container is 5.16 GB. At 512px a guided batch-1 run needs a little over 7 GB of +GPU memory in eager Torch (about 1.5x diffusers' eager time); at 768px the self-attention +runs over 9216 tokens, so the guided run needs roughly 12 GB, or the CPU. + +## Loading Fine-tuned Weights + +Any repo laid out like the hosted ones (`zm_config.json` declaring +`StableDiffusion2Model`, the weights, `tokenizer.json`) loads with +`from_weights("/")`. The `hf:` prefix raises for diffusion models: convert a +diffusers-format SD 2 checkpoint once with +`zeromodels/models/stable_diffusion_2/convert_stable_diffusion_2_diffusers_to_keras.py` +(`transfer_stable_diffusion_2(repo)`, `pip install zeromodels[conversion]`) and host the result. diff --git a/docs/stable_diffusion_3.md b/docs/stable_diffusion_3.md new file mode 100644 index 00000000..bf699bbe --- /dev/null +++ b/docs/stable_diffusion_3.md @@ -0,0 +1,268 @@ +# Stable Diffusion 3 + +
+Weights: pretrained Keras weights live on Hugging Face under +zeromodels/<variant> +(each repo carries zm_config.json + the float16 weights + +tokenizer.json / tokenizer_3.json). Load with +from_weights("zeromodels/<variant>"). +
+ +Stable Diffusion 3, ported to pure Keras 3: a rectified-flow model whose denoiser is the +**MMDiT**, a multimodal diffusion transformer in which the latent tokens and the text tokens +attend to each other jointly, conditioned on the timestep and the pooled text embeddings +through adaptive layer norms. It keeps the SD family's shape (one container per repo, +`generate` from `BaseDiffusion`, the schedulers, both data formats) with new parts: + +- **MMDiT** (`SD3Transformer2DModel`): the 16-channel latent is patchified (2x2) into 1536-d + tokens with a sinusoidal position table cropped from a 192x192 grid; 24 joint blocks of + 24 heads, each with AdaLN-Zero modulation of both streams, joint attention over the + concatenated latent + text tokens, and gated GELU feed-forwards; a final adaptive norm + and projection unpatchify the velocity. 2B parameters. +- **Three text encoders**: CLIP ViT-L/14 (768-d) and OpenCLIP ViT-bigG/14 (1280-d), each + with a projection, and **T5-XXL** (4.7B, 4096-d, 256 tokens). The two CLIP penultimate + states are concatenated (2048-d), zero-padded to 4096 and followed along the sequence by + the T5 states (77 + 256 tokens); the projected pooled CLIP states (2048-d) feed the + conditioning embedding. +- **T5 is optional and separate**: the container holds the MMDiT, the VAE and the two CLIP + towers (2.9B, 5.8 GB float16); the T5-XXL encoder is hosted once for every SD 3 / 3.5 + checkpoint (`zeromodels/t5-v1_1-xxl-encoder`, an + [`SD3T5EncoderModel`](#sd3t5encodermodel), 9.5 GB float16) and attached with + `text_encoder_3=`. Without it the T5 features are zeros, SD 3's documented + memory-saving mode (prompt adherence drops, the images stay good). +- **VAE**: 16 latent channels, no quant convolutions, `scaling_factor` 1.5305 and + `shift_factor` 0.0609 (`z = (x - shift) * scale`), built in float32 (`force_upcast`). +- **Sampler**: `FlowMatchEulerDiscreteScheduler` with `shift` 3.0 (rectified flow: the + timesteps are `sigma * 1000`, a step is `x + (sigma_next - sigma) * v`); 28 steps at + guidance 7.0 by default. No negative prompt means the encoded empty prompt. +- **One tokenizer call**: `StableDiffusion3Tokenizer` runs the CLIP BPE (the second + tower's `!`-padded ids are derived from the mask, as SDXL) and the T5 SentencePiece + tokenizer, returning `{"input_ids", "attention_mask", "input_ids_3"}`. +- **float16 weights**, the release precision; `from_weights` builds the model in float16 + by default (`load_dtype="float32"` for float32). + +The weights are converted once, offline, and hosted: on-the-fly `hf:` conversion is +deliberately **not supported** for diffusion models. The upstream repo is gated behind the +Stability AI Non-Commercial Research Community License. + +Links: + +- Paper: [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis (arXiv:2403.03206)](https://arxiv.org/abs/2403.03206) +- Reference implementation: [diffusers `StableDiffusion3Pipeline`](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_3) +- License: [Stability AI Non-Commercial Research Community License](https://huggingface.co/stabilityai/stable-diffusion-3-medium-diffusers/blob/main/LICENSE.md) + +See also [stable_diffusion_3_5.md](stable_diffusion_3_5.md), [stable_diffusion_xl.md](stable_diffusion_xl.md), [clip.md](clip.md). + +## Variants + +| Variant | Hub | Resolution | Sampler | License | +|---|---|---|---|---| +| `stable-diffusion-3-medium` | [`zeromodels/stable-diffusion-3-medium`](https://huggingface.co/zeromodels/stable-diffusion-3-medium) | 1024 | flow-match Euler (shift 3), 28 steps, guidance 7.0 | Stability AI Non-Commercial Research Community | +| `t5-v1_1-xxl-encoder` | [`zeromodels/t5-v1_1-xxl-encoder`](https://huggingface.co/zeromodels/t5-v1_1-xxl-encoder) | text encoder (shared with SD 3.5) | | Apache 2.0 (Google T5 v1.1) | + +## API + +Configs are typed: `StableDiffusion3Config` (composite, `model_type` `"stable_diffusion_3"`) +over `StableDiffusion3TransformerConfig`, `StableDiffusion3VAEConfig`, +`StableDiffusion3TextConfig` (CLIP ViT-L/14 with its 768-d projection) and +`StableDiffusionXLTextConfig2` (OpenCLIP ViT-bigG/14 with its 1280-d projection), plus +the scheduler config, the CLIP and T5 token ids and `max_sequence_length` (256). Flat +constructor with `transformer_` / `vae_` / `text_` / `text_2_` prefixes. + +### `StableDiffusion3TextToImage` + +`StableDiffusionTextToImage`'s `generate` over the SD 3 graph: + +```python +generate( + input_ids, + attention_mask=None, + negative_input_ids=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, + image=None, + strength=None, + denoising_start=None, + denoising_end=None, + output_type="image", + input_ids_3=None, + negative_input_ids_3=None, +) +``` + +| Argument | Description | +|---|---| +| `input_ids`, `attention_mask`, `input_ids_3` | `**tokenizer(prompts)`: the CLIP ids and mask (the second tower's ids are derived from the mask) and the T5 ids (`(batch, 256)`). Without `input_ids_3`, or without an attached T5, the T5 features are zeros. | +| `negative_input_ids`, `negative_input_ids_3` | The tokenized negative prompt (`neg = tokenizer("..."); negative_input_ids=neg["input_ids"], negative_input_ids_3=neg["input_ids_3"]`); the empty prompt when omitted. | +| `num_inference_steps`, `guidance_scale`, `seed`, `latents`, `image`, `strength`, ... | As for [`BaseDiffusion.generate`](main_classes.md#basediffusion); defaults from the repo's `generate_args` (28 steps, guidance 7.0). `latents` is `(batch, 128, 128, 16)` at 1024px. | + +Returns `(batch, H, W, 3)` uint8 images. + +`from_weights` takes one extra argument, `text_encoder_3`: a hosted `SD3T5EncoderModel` +repo id (loaded at the same `load_dtype`) or a built `SD3T5EncoderModel`; +`sd.text_encoder_3` can also be assigned later (or set to `None`). It lives outside the +container's weights. + +### `StableDiffusion3Model` + +The container: four disconnected paths in one functional model, `.transformer` / `.vae` / +`.text_encoder` / `.text_encoder_2`. + +| Inputs | Outputs | +|---|---| +| `sample` (B, 128, 128, 16), `timestep` (B,), `encoder_hidden_states` (B, 333, 4096), `pooled_projections` (B, 2048) | `noise_pred` (B, 128, 128, 16) | +| `image` (B, 1024, 1024, 3), `latent` (B, 128, 128, 16) | `moments` (B, 128, 128, 32), `image` (B, 1024, 1024, 3) | +| `token_ids` (B, 77), `token_ids_2` (B, 77), `padding_mask` (B, 77) | `prompt_embeds` (B, 77, 2048), `pooled_prompt_embeds` (B, 2048) | + +### `SD3Transformer2DModel` + +The MMDiT on its own (`{"sample", "timestep", "encoder_hidden_states", +"pooled_projections"} -> {"sample"}`), built from `StableDiffusion3TransformerConfig` +(`sample_size` 128, `patch_size` 2, `num_layers` 24, `num_attention_heads` 24 x +`attention_head_dim` 64, `joint_attention_dim` 4096, `caption_projection_dim` 1536, +`pooled_projection_dim` 2048, `pos_embed_max_size` 192, `qk_norm`, +`dual_attention_layers`). Its layers (`StableDiffusion3PatchEmbed`, +`StableDiffusion3AdaLayerNorm`, `StableDiffusion3JointAttention`, +`StableDiffusion3JointTransformerBlock`, `StableDiffusion3GELUFeedForward`) are named +with the diffusers module paths. + +The joint attention runs over 4429 tokens at 1024px, so `StableDiffusion3JointAttention` +defaults to the `"fused"` attention implementation (`keras.ops.dot_product_attention`, +the backend's own kernel: torch's flash / memory-efficient kernels never materialize the +4429 x 4429 logits; see [`fused_attention`](main_classes.md#fused_attention)) instead of +the library-wide `"sdpa"` math. `from_weights(..., attn_implementation="sdpa")` switches +it back. + +### `SD3T5EncoderModel` + +The third text encoder, the encoder half of T5 v1.1 XXL (transformers `T5EncoderModel` on +`google/t5-v1_1-xxl`), as its own model in this family: `{"input_ids", "attention_mask"}` +(`(batch, 256)` int32) `-> {"last_hidden_state": (batch, 256, 4096)}`, built from +`StableDiffusion3T5EncoderConfig` (`vocab_size` 32128, `embed_dim` 4096, `key_value_dim` +64, `mlp_dim` 10240, `num_layers` 24, `num_heads` 64, `relative_attention_num_buckets` 32, +`relative_attention_max_distance` 128, `layer_norm_eps` 1e-6; 4.7B parameters). The blocks +reuse the [T5](t5.md) family's attention, RMSNorm and relative position bias, and add the +gated-GELU feed-forward of T5 v1.1 (`wo(gelu_tanh(wi_0(x)) * wi_1(x))`, +`StableDiffusion3T5GatedFeedForward`), which the original T5 does not have. Loads on its own with +`SD3T5EncoderModel.from_weights("zeromodels/t5-v1_1-xxl-encoder")` (`load_dtype`, +`quantization="int8"` as needed); its inputs are the `input_ids_3` of the SD 3 tokenizer. + +## Preprocessing + +### `StableDiffusion3Tokenizer` + +```python +StableDiffusion3Tokenizer( + hf_id=None, + tokenizer_file=None, + tokenizer_file_3=None, + max_seq_len=77, + max_sequence_length=256, +) +``` + +Returns `{"input_ids", "attention_mask", "input_ids_3"}`; pass all three to `generate`. + +## End-to-end example + +```python +import os + +os.environ["KERAS_BACKEND"] = "torch" # or "jax" / "tensorflow" + +from PIL import Image +from zeromodels.models.stable_diffusion_3 import ( + StableDiffusion3TextToImage, + StableDiffusion3Tokenizer, +) + +model = StableDiffusion3TextToImage.from_weights( + "zeromodels/stable-diffusion-3-medium", + text_encoder_3="zeromodels/t5-v1_1-xxl-encoder", # omit on small machines: T5 features zeroed +) +tokenizer = StableDiffusion3Tokenizer.from_weights( + "zeromodels/stable-diffusion-3-medium" +) + +inputs = tokenizer( + "a fluffy corgi wearing round sunglasses, sitting on a surfboard at a sunny beach, photo" +) +images = model.generate(**inputs, num_inference_steps=28, guidance_scale=7.0, seed=4) + +Image.fromarray(images[0]).save("corgi.png") # (1024, 1024, 3) uint8 +``` + +Stable Diffusion 3 medium: a corgi with sunglasses on a surfboard, 1024px + +### The T5 encoder on the CPU + +The 9.5 GB T5-XXL rarely fits next to the MMDiT on a consumer GPU. Encode the prompt +separately (an `SD3T5EncoderModel` in another process, or int8) and hand the ids to a +model without T5 for the CLIP part only, or attach it in full on a large machine. The +`encode_prompt` hook returns the conditioning dict the transformer takes, so any +precomputed T5 features can be substituted there. + +### Verified against diffusers + +The components of `stable-diffusion-3-medium` were checked one by one against diffusers +0.39 / transformers 5 in float32 on the same inputs (the converted float16 weights +computed in float32): + +``` +text_encoder penultimate max|d|=6.1e-05 (values up to 854) +text_encoder_2 penultimate max|d|=3.1e-05 +pooled (L + G, projected) max|d|=3.8e-06 +t5-v1_1-xxl encoder (256 tok) max|d|=1.1e-04 (values up to 6.3) +mmdit velocity (t=1000) max|d|=9.9e-06 +vae decode (512px) max|d|=6.3e-05 +vae encode max|d|=1.4e-05 +``` + +The pipeline (flow-match Euler with shift 3, classifier-free guidance with the encoded +empty prompt, the three tokenizers, batched negative prompts, the no-T5 mode) was checked +against `StableDiffusion3Pipeline` on a small random-weight SD 3, where the generated images +are pixel-identical on torch and jax. End to end, `stable-diffusion-3-medium` without the +T5 encoder (28 steps, guidance 7, the same initial latent) reproduces the diffusers +pipeline in float32 at 512px: 92.8% identical pixels, max |d| 2 / 255, PSNR 63.9 dB (the +final latent differs by at most 2.5e-2 with values up to 5.1). In float16 on the GPU the +two implementations round differently through the 28 guided steps: the pictures agree in +composition but not pixel for pixel (PSNR 26 dB at 1024px, 29 dB at 768px, float16 MMDiT +on both sides), the usual float16 sampling noise. + +## Data Format + +Both `channels_last` and `channels_first` are supported; `generate` always returns +`(batch, H, W, 3)` uint8. + +## Memory and speed + +The container is 5.8 GB in float16 (the T5 encoder adds 9.5 GB). Measured on the torch +backend with the default float16 load (RTX 4060 Laptop, 8 GB), T5 features precomputed: + +| Resolution | Denoising step (guided, batch 1) | VAE decode (float32) | +|---|---|---| +| 512px | peak 6.1 GB, 0.5 s/step | peak 7.6 GB, 4 s | +| 768px | peak 6.4 GB, 0.8 s/step | peak 9.5 GB | +| 1024px | peak 6.9 GB, 1.0 s/step | peak 12.4 GB | + +The MMDiT itself fits an 8 GB GPU at the native 1024px thanks to the fused attention +kernel (the `"sdpa"` math needs 3.8 GB per block for the 4429 x 4429 float32 logits and +runs out of memory). The VAE decoder is the SDXL one (16 latent channels aside) and Keras' +functional executor keeps its feature maps alive, see the +[SDXL notes](stable_diffusion_xl.md#memory-and-speed): on 8 GB, denoise on the GPU +(`output_type="latent"`) and decode on the CPU (75 to 90 s at 1024px). The T5-XXL encoder +encodes a prompt and its negative on the CPU in about 10 to 25 s (float32 compute over +the float16 weights); a 24 GB GPU holds everything. On JAX and TensorFlow the denoiser +step is compiled once per run. + +## Loading Fine-tuned Weights + +Any repo laid out like the hosted ones (`zm_config.json` declaring +`StableDiffusion3Model`, the weights, the two tokenizer files) loads with +`from_weights("/")`. The `hf:` prefix raises for diffusion models: convert a +diffusers-format SD 3 checkpoint once with +`zeromodels/models/stable_diffusion_3/convert_stable_diffusion_3_diffusers_to_keras.py` +(`transfer_stable_diffusion_3(repo)`; `transfer_t5_encoder(repo)` for the T5 tower) and host the +result. diff --git a/docs/stable_diffusion_3_5.md b/docs/stable_diffusion_3_5.md new file mode 100644 index 00000000..5a2e843e --- /dev/null +++ b/docs/stable_diffusion_3_5.md @@ -0,0 +1,128 @@ +# Stable Diffusion 3.5 + +
+Weights: pretrained Keras weights live on Hugging Face under +zeromodels/<variant> +(each repo carries zm_config.json + the float16 weights + +tokenizer.json / tokenizer_3.json). Load with +from_weights("zeromodels/<variant>"). +
+ +Stable Diffusion 3.5, ported to pure Keras 3: the [Stable Diffusion 3](stable_diffusion_3.md) +architecture, retrained and scaled. It reuses the SD 3 family end to end +(`StableDiffusion3_5Model` is `StableDiffusion3Model`, `StableDiffusion3_5TextToImage` is +`StableDiffusion3TextToImage`, each with the SD 3.5 configuration), so everything on that +page applies: the MMDiT, the three text encoders with the separately hosted T5-XXL, the +16-channel VAE, the flow-match sampler, one tokenizer call. What changes is the config: + +- **QK normalization**: every attention block RMS-normalizes its queries and keys per head + (`qk_norm="rms_norm"`), the training stabilizer of SD 3.5. +- **Large** (8B): 38 blocks of 38 heads (2432-d tokens), 28 steps at guidance 3.5. +- **Large Turbo**: the large model distilled with Adversarial Diffusion Distillation to 4 + steps without guidance (`guidance_scale=0.0`). +- **Medium** (2.5B, MMDiT-X): 24 blocks of 24 heads with **dual attention** in blocks 0 to + 12 (a second, latent-only attention with its own modulation next to the joint one) and a + 384x384 position grid (up to 2 MP); 40 steps at guidance 4.5. + +The weights are converted once, offline, and hosted: on-the-fly `hf:` conversion is +deliberately **not supported** for diffusion models. The upstream repos are gated behind +the Stability AI Community License. + +Links: + +- Paper: [Scaling Rectified Flow Transformers for High-Resolution Image Synthesis (arXiv:2403.03206)](https://arxiv.org/abs/2403.03206) +- Reference implementation: [diffusers `StableDiffusion3Pipeline`](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_3) +- License: [Stability AI Community License](https://huggingface.co/stabilityai/stable-diffusion-3.5-large/blob/main/LICENSE.md) + +See also [stable_diffusion_3.md](stable_diffusion_3.md). + +## Variants + +Each repo is one container: MMDiT + VAE (84M) + CLIP ViT-L/14 (123M) + OpenCLIP +ViT-bigG/14 (695M); the T5-XXL encoder (`zeromodels/t5-v1_1-xxl-encoder`, the SD 3 +family's `SD3T5EncoderModel`, 9.5 GB) is shared and optional. + +| Variant | Hub | MMDiT | Sampler | Size | +|---|---|---|---|---| +| `stable-diffusion-3.5-large` | [`zeromodels/stable-diffusion-3.5-large`](https://huggingface.co/zeromodels/stable-diffusion-3.5-large) | 38 x 38 heads, 8B | flow-match Euler (shift 3), 28 steps, guidance 3.5 | 17.8 GB float16 | +| `stable-diffusion-3.5-large-turbo` | [`zeromodels/stable-diffusion-3.5-large-turbo`](https://huggingface.co/zeromodels/stable-diffusion-3.5-large-turbo) | 38 x 38 heads, 8B | 4 steps, no guidance | 17.8 GB float16 | +| `stable-diffusion-3.5-medium` | [`zeromodels/stable-diffusion-3.5-medium`](https://huggingface.co/zeromodels/stable-diffusion-3.5-medium) | 24 x 24 heads + dual attention, 2.5B | flow-match Euler (shift 3), 40 steps, guidance 4.5 | 6.8 GB float16 | + +## API + +`StableDiffusion3_5Config` (`model_type` `"stable_diffusion_3_5"`) is +`StableDiffusion3Config` over `StableDiffusion3_5TransformerConfig` (the SD 3.5 large +defaults: `num_layers` 38, `num_attention_heads` 38, `caption_projection_dim` 2432, +`qk_norm="rms_norm"`; medium's repo carries `dual_attention_layers=(0, ..., 12)` and +`pos_embed_max_size=384`). `StableDiffusion3_5TextToImage.generate` and +`StableDiffusion3_5Tokenizer` are the SD 3 ones; see +[Stable Diffusion 3](stable_diffusion_3.md#api). + +## End-to-end example + +```python +import os + +os.environ["KERAS_BACKEND"] = "torch" # or "jax" / "tensorflow" + +from PIL import Image +from zeromodels.models.stable_diffusion_3_5 import ( + StableDiffusion3_5TextToImage, + StableDiffusion3_5Tokenizer, +) + +model = StableDiffusion3_5TextToImage.from_weights( + "zeromodels/stable-diffusion-3.5-medium", + text_encoder_3="zeromodels/t5-v1_1-xxl-encoder", # optional +) +tokenizer = StableDiffusion3_5Tokenizer.from_weights( + "zeromodels/stable-diffusion-3.5-medium" +) + +inputs = tokenizer( + "a whimsical treehouse village at twilight, lanterns glowing, watercolor illustration" +) +images = model.generate(**inputs, num_inference_steps=40, guidance_scale=4.5, seed=5) + +Image.fromarray(images[0]).save("treehouse.png") # (1024, 1024, 3) uint8 +``` + +Stable Diffusion 3.5 medium: a treehouse village at twilight, 1024px + +### Large Turbo + +```python +model = StableDiffusion3_5TextToImage.from_weights( + "zeromodels/stable-diffusion-3.5-large-turbo" +) +images = model.generate( + **tokenizer("a red panda reading a book") +) # 4 steps, guidance 0 (the repo defaults) +``` + +### Verified against diffusers + +The components of `stable-diffusion-3.5-medium` (the MMDiT-X with dual attention and RMS +qk norms) were checked against diffusers 0.39 in float32 on the same inputs (the converted +float16 weights computed in float32): MMDiT velocity max|d| 1.6e-5 (values up to 5.2), the +text encoders and the VAE as for [SD 3](stable_diffusion_3.md#verified-against-diffusers). +The SD 3.5 block variants (dual attention, `qk_norm`) are also covered by the small +random-weight pipeline test, pixel-identical to `StableDiffusion3Pipeline` on torch and +jax, and by the SD 3 float32 end-to-end check (92.8% identical pixels, PSNR 63.9 dB at +512px), which the two families share down to the sampler. The picture above is the +1024px float16 output of `stable-diffusion-3.5-medium` with the T5 encoder. + +## Memory and speed + +The medium container is 6.8 GB in float16, the large ones 17.8 GB (a 24 GB GPU); the +T5-XXL encoder adds 9.5 GB. See [Stable Diffusion 3](stable_diffusion_3.md#memory-and-speed): +the medium MMDiT-X at 1024px peaks at 7.5 GB during the guided denoising step on the torch +backend (fused attention, float16), which an 8 GB GPU runs with the VAE decode on the CPU. + +## Loading Fine-tuned Weights + +Any repo laid out like the hosted ones (`zm_config.json` declaring +`StableDiffusion3_5Model`, the weights, the two tokenizer files) loads with +`from_weights("/")`. Convert a diffusers-format checkpoint once with +`zeromodels/models/stable_diffusion_3_5/convert_stable_diffusion_3_5_diffusers_to_keras.py` +and host the result. diff --git a/docs/stable_diffusion_xl.md b/docs/stable_diffusion_xl.md new file mode 100644 index 00000000..c62ad479 --- /dev/null +++ b/docs/stable_diffusion_xl.md @@ -0,0 +1,385 @@ +# Stable Diffusion XL + +
+Weights: pretrained Keras weights live on Hugging Face under +zeromodels/<variant> +(each repo carries zm_config.json + the float16 weights + +tokenizer.json). Load with from_weights("zeromodels/<variant>"). +
+ +Stable Diffusion XL, ported to pure Keras 3: the [Stable Diffusion](stable_diffusion.md) +latent-diffusion recipe scaled up to 1024px with a 2.6B-parameter UNet and two text +encoders, plus the **refiner**, an image-to-image expert that finishes the base model's +latents, and **SDXL-Turbo**, the base distilled to a few steps. It subclasses the SD 1.x +family (`StableDiffusionXLModel` is a `StableDiffusionModel` with a fourth component, +`StableDiffusionXLTextToImage` is `StableDiffusionTextToImage` with the SDXL conditioning +hooks; the refiner classes subclass those), so everything on that page applies: one +container per repo, `generate` from `BaseDiffusion`, the schedulers, both data formats. +What changes: + +- **Two text encoders**: CLIP ViT-L/14's text tower (768-d, the SD 1.x one) and OpenCLIP + ViT-bigG/14's (1280-d, 32 layers, `gelu`, with a 1280-d `text_projection`). The prompt + goes through both; their **penultimate** hidden states (no final LayerNorm) are + concatenated into the UNet's 2048-d cross-attention context, and the second tower's + projected pooled (EOT) state is the `text_embeds` micro-conditioning. +- **One tokenizer call**: both towers use the same CLIP BPE and differ only in the pad + token (`<|endoftext|>` vs `!`), so `StableDiffusionXLTokenizer` is the first one and the + model derives the second tower's ids from the returned `attention_mask`. +- **Micro-conditioning**: the UNet's timestep embedding also receives the pooled text + embedding and six size / crop values (`original_size`, `crops_coords_top_left`, + `target_size`, each `(height, width)`), sinusoidally embedded and projected by the + `add_embedding` MLP (`addition_embed_type="text_time"`). +- **UNet**: three levels (320, 640, 1280) with no attention at the first, (1, 2, 10) + transformer blocks per Transformer2D, (5, 10, 20) heads (64 wide), linear token + projection, built at a 128x128 latent (1024px). +- **VAE**: the SDXL autoencoder (`scaling_factor` 0.13025), built in **float32 whatever + the load dtype** (`force_upcast`: it overflows in float16, as in diffusers). +- **Classifier-free guidance** with no negative prompt uses **zero** embeddings + (`force_zeros_for_empty_prompt`), not the encoded empty prompt. Default guidance is 5.0. +- **Scheduler**: Euler with `leading` timestep spacing (`steps_offset` 1) for the base + models and the refiners, ancestral Euler with `trailing` spacing for SDXL-Turbo; all + come from the repo's `scheduler_config`. +- **Refiner**: a second model (`StableDiffusionXLRefinerModel` / + `StableDiffusionXLRefinerImageToImage`) with the OpenCLIP tower alone (1280-d + context), a four-level UNet ((384, 768, 1536, 1536), attention at the middle levels, 4 + blocks per Transformer2D, 2.3B parameters) and an **aesthetic score** as the fifth + micro-conditioning id. It refines a latent the base left partially denoised + (`denoising_start`) or any image (`strength`); its empty negative prompt is encoded, not + zeroed. +- **float16 weights**: SDXL was released in float16, so the repos store that (the VAE in + float32) and `from_weights` builds the model in float16 by default (6.6 GB of weights). + Pass `load_dtype="float32"` for a float32 model (13.9 GB). + +The weights are converted once, offline, and hosted: on-the-fly `hf:` conversion is +deliberately **not supported** for diffusion models. + +Links: + +- Paper: [SDXL: Improving Latent Diffusion Models for High-Resolution Image Synthesis (arXiv:2307.01952)](https://arxiv.org/abs/2307.01952) +- Reference implementation: [diffusers `StableDiffusionXLPipeline`](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl) +- Licenses: [CreativeML Open RAIL++-M](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/blob/main/LICENSE.md) + (1.0 base and refiner), [Stability AI Non-Commercial Research Community License](https://huggingface.co/stabilityai/sdxl-turbo/blob/main/LICENSE.md) + (Turbo); the 0.9 research previews are not hosted + +See also [stable_diffusion.md](stable_diffusion.md), [stable_diffusion_2.md](stable_diffusion_2.md), [clip.md](clip.md). + +## Variants + +Preconverted weights are hosted under `zeromodels/`. Load with +`from_weights("zeromodels/")`. A base repo is one container: UNet (2.57B) + VAE +(84M) + CLIP ViT-L/14 text encoder (123M) + OpenCLIP ViT-bigG/14 text encoder (695M), +3.47B parameters, 6.6 GB in float16 (two `model.weights.json` shards); a refiner repo is +UNet (2.26B) + VAE + OpenCLIP text encoder, 3.04B parameters, 5.8 GB. + +| Variant | Hub | Class | Resolution | Sampler | License | +|---|---|---|---|---|---| +| `stable-diffusion-xl-base-1.0` | [`zeromodels/stable-diffusion-xl-base-1.0`](https://huggingface.co/zeromodels/stable-diffusion-xl-base-1.0) | `StableDiffusionXLTextToImage` | 1024 | Euler, 50 steps, guidance 5.0 | CreativeML Open RAIL++-M | +| `stable-diffusion-xl-refiner-1.0` | [`zeromodels/stable-diffusion-xl-refiner-1.0`](https://huggingface.co/zeromodels/stable-diffusion-xl-refiner-1.0) | `StableDiffusionXLRefinerImageToImage` | 1024 | Euler, 50 steps, guidance 5.0, strength 0.3 | CreativeML Open RAIL++-M | +| `sdxl-turbo` | [`zeromodels/sdxl-turbo`](https://huggingface.co/zeromodels/sdxl-turbo) | `StableDiffusionXLTextToImage` | 512 | Euler ancestral, 1 step, no guidance | Stability AI Non-Commercial Research Community | + +The SDXL 0.9 research previews (`stabilityai/stable-diffusion-xl-base-0.9` and +`-refiner-0.9`, the same architecture with different weights, gated behind the research +license) are not hosted; with an authorized token the converter turns them into the same +containers (`transfer_stable_diffusion_xl(repo, token=...)`, see +[Loading Fine-tuned Weights](#loading-fine-tuned-weights)). SDXL-Turbo is a +research-only, non-commercial release; its repo defaults (`generate_args`) are 1 step and +`guidance_scale=0.0`, which is how it was trained (do not add guidance). + +## API + +Configs are typed: `StableDiffusionXLConfig` (composite, `model_type` +`"stable_diffusion_xl"`) over `StableDiffusionXLUNetConfig`, `StableDiffusionXLVAEConfig`, +`StableDiffusionTextConfig` (the CLIP ViT-L/14 tower, as SD 1.x) and +`StableDiffusionXLTextConfig2` (the OpenCLIP ViT-bigG/14 tower, with `projection_dim` and +its own `hidden_act`), plus the scheduler config, the token ids (`pad_token_id_2` for the +second tokenizer's `!`), `force_zeros_for_empty_prompt` and `requires_aesthetics_score`. +Flat constructor with `unet_` / `vae_` / `text_` / `text_2_` prefixes. +`StableDiffusionXLRefinerConfig` (`model_type` `"stable_diffusion_xl_refiner"`) is the +same with `text_config=None`, the `StableDiffusionXLRefinerUNetConfig` denoiser, +`requires_aesthetics_score=True` and `force_zeros_for_empty_prompt=False`. + +### `StableDiffusionXLTextToImage` + +`StableDiffusionTextToImage` over the SDXL graph; `generate` gains the micro-conditioning +arguments: + +```python +generate( + input_ids, + attention_mask=None, + negative_input_ids=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, + image=None, + strength=None, + denoising_start=None, + denoising_end=None, + output_type="image", + original_size=None, + crops_coords_top_left=(0, 0), + target_size=None, + aesthetic_score=6.0, + negative_original_size=None, + negative_crops_coords_top_left=None, + negative_target_size=None, + negative_aesthetic_score=2.5, +) +``` + +| Argument | Description | +|---|---| +| `input_ids`, `attention_mask` | `**tokenizer(prompts)`; the mask tells the model which positions are padding, so the second tower's `!`-padded ids can be derived (without it, everything after the first `<|endoftext|>` counts as padding). | +| `negative_input_ids` | Tokenized negative prompt, one row per prompt or a single row shared by the batch. Without it the unconditional branch is zero embeddings (`force_zeros_for_empty_prompt`), as in diffusers. | +| `num_inference_steps`, `guidance_scale`, `seed`, `latents` | As for [Stable Diffusion](stable_diffusion.md#stablediffusiontexttoimage); defaults from the repo's `generate_args` (50 steps, guidance 5.0 for the base model). | +| `image`, `strength`, `denoising_start`, `denoising_end`, `output_type` | The image-to-image and ensemble controls of [`BaseDiffusion.generate`](main_classes.md#basediffusion): refine an `image` (noised to `strength`), stop early (`denoising_end`) and hand the latent over (`output_type="latent"`), or resume one (`denoising_start`). | +| `original_size`, `crops_coords_top_left`, `target_size` | SDXL's size / crop conditioning, `(height, width)` in pixels; the sizes default to the generated image's size, the crop to `(0, 0)`. | +| `aesthetic_score` | The refiner's fifth conditioning value (6.0), in place of `target_size`; ignored by the base models. | +| `negative_*` | The same for the negative branch; each defaults to its positive value, except `negative_aesthetic_score` (2.5). | + +Returns `(batch, H, W, 3)` uint8 images (or the latent); `latents` is +`(batch, 128, 128, 4)` at 1024px (`channels_first`: channels second). + +The hooks, for anyone composing their own loop: `encode_prompt(input_ids, attention_mask, +original_size, crops_coords_top_left, target_size)` returns the dict the UNet takes +(`encoder_hidden_states` (batch, 77, 2048), `text_embeds` (batch, 1280), `time_ids` +(batch, 6)); `encode_negative_prompt` zeroes the first two when no negative ids are given; +`predict_noise(latents, timesteps, embeddings)` runs the UNet. + +### `StableDiffusionXLRefinerImageToImage` + +The refiner task: the same class over the refiner container (`text_encoder_2` alone, +`aesthetic_score` in the conditioning, `generate_args` with `strength` 0.3). Its +`generate` is the one above; with no `image` / `latents` it runs text-to-image, which the +refiner was not trained for. The two intended uses: + +```python +base = StableDiffusionXLTextToImage.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0" +) +refiner = StableDiffusionXLRefinerImageToImage.from_weights( + "zeromodels/stable-diffusion-xl-refiner-1.0" +) +tokenizer = StableDiffusionXLTokenizer.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0" +) +inputs = tokenizer("a lighthouse on a rocky coast at dawn, dramatic clouds, cinematic") + +# ensemble of experts: the base denoises 80% of the schedule, the refiner the rest +latent = base.generate( + **inputs, num_inference_steps=50, denoising_end=0.8, output_type="latent" +) +images = refiner.generate( + **inputs, latents=latent, num_inference_steps=50, denoising_start=0.8 +) + +# image-to-image: refine any image (uint8 or [0, 1] float, (batch, H, W, 3)) +images = refiner.generate(**inputs, image=images, strength=0.3, seed=0) +``` + +### `StableDiffusionXLModel` + +The container: four disconnected paths in one functional model, `.unet` / `.vae` / +`.text_encoder` / `.text_encoder_2`. + +| Inputs | Outputs | +|---|---| +| `sample` (B, 128, 128, 4), `timestep` (B,), `encoder_hidden_states` (B, 77, 2048), `text_embeds` (B, 1280), `time_ids` (B, 6) | `noise_pred` (B, 128, 128, 4) | +| `image` (B, 1024, 1024, 3), `latent` (B, 128, 128, 4) | `moments` (B, 128, 128, 8), `image` (B, 1024, 1024, 3) | +| `token_ids` (B, 77), `token_ids_2` (B, 77), `padding_mask` (B, 77) | `prompt_embeds` (B, 77, 2048), `pooled_prompt_embeds` (B, 1280) | + +Each text tower is a functional CLIP text model (the `clip` family's layers) whose outputs +are `penultimate_hidden_state`, `last_hidden_state`, `pooler_output` and, for the second, +`text_embeds`. The UNet is the SD 1.x `UNet2DConditionModel` with the SDXL config +(`transformer_layers_per_block=(1, 2, 10)`, `addition_embed_type="text_time"`); the VAE +is `AutoencoderKL` with `force_upcast=True`. `StableDiffusionXLRefinerModel` is the same +container without the first tower (`token_ids_2` + `padding_mask` in, `prompt_embeds` +(B, 77, 1280) out) and with `time_ids` (B, 5). + +## Preprocessing + +### `StableDiffusionXLTokenizer` + +`StableDiffusionTokenizer` (CLIP BPE, `<|startoftext|>` / `<|endoftext|>` framing, +`<|endoftext|>`-padded to 77 tokens). Returns `{"input_ids", "attention_mask"}`; pass both +to `generate`. + +```python +StableDiffusionXLTokenizer( + hf_id=None, tokenizer_file=None, max_seq_len=77, pad_token="<|endoftext|>" +) +``` + +A literal `!` in a prompt is encoded as OpenCLIP encodes it (the BPE token `!`, id +256) in the second tower's derived ids; diffusers' `tokenizer_2`, whose pad token is `!`, +turns it into the pad id 0 instead. + +## End-to-end example + +```python +import os + +os.environ["KERAS_BACKEND"] = "torch" # or "jax" / "tensorflow" + +from PIL import Image +from zeromodels.models.stable_diffusion_xl import ( + StableDiffusionXLTextToImage, + StableDiffusionXLTokenizer, +) + +model = StableDiffusionXLTextToImage.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0" +) +tokenizer = StableDiffusionXLTokenizer.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0" +) + +inputs = tokenizer( + "a lighthouse on a rocky coast at dawn, dramatic clouds, cinematic, highly detailed" +) +images = model.generate(**inputs, num_inference_steps=30, guidance_scale=5.0, seed=3) + +Image.fromarray(images[0]).save("lighthouse.png") # (1024, 1024, 3) uint8 +``` + +Stable Diffusion XL base 1.0: a lighthouse on a rocky coast at dawn, 1024px + +### Base + refiner + +The SDXL "ensemble of experts": the base runs the first 80% of the schedule and hands its +latent to the refiner, which finishes it. Same prompt and initial latent as above: + +```python +refiner = StableDiffusionXLRefinerImageToImage.from_weights( + "zeromodels/stable-diffusion-xl-refiner-1.0" +) +latent = model.generate( + **inputs, num_inference_steps=30, denoising_end=0.8, output_type="latent", seed=3 +) +images = refiner.generate( + **inputs, latents=latent, num_inference_steps=30, denoising_start=0.8 +) +``` + +Stable Diffusion XL base + refiner: the lighthouse, refined + +### SDXL-Turbo + +One step, no guidance, 512px (the repo's defaults): + +```python +model = StableDiffusionXLTextToImage.from_weights("zeromodels/sdxl-turbo") +tokenizer = StableDiffusionXLTokenizer.from_weights("zeromodels/sdxl-turbo") +images = model.generate( + **tokenizer( + "a cinematic shot of a baby raccoon wearing an intricate italian priest robe" + ) +) +``` + +SDXL-Turbo, one step: a baby raccoon in an italian priest robe, 512px + +### Negative prompts, micro-conditioning + +```python +inputs = tokenizer(["a red bicycle", "a blue bicycle"]) +negative = tokenizer("blurry, low quality") # one row, shared by the batch +images = model.generate( + **inputs, + negative_input_ids=negative["input_ids"], + original_size=(2048, 2048), # "a crop of a larger image": sharper detail + crops_coords_top_left=(0, 0), + target_size=(1024, 1024), + seed=0, +) +``` + +### Other resolutions + +The graphs are built for the repo's resolution; pass overrides to build for another +(the weights are resolution-independent; multiples of 64px): + +```python +model = StableDiffusionXLTextToImage.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0", + unet_sample_size=(96, 128), + vae_sample_size=(768, 1024), +) +images = model.generate(**tokenizer("a wide mountain valley")) # (1, 768, 1024, 3) +``` + +### float32 + +```python +model = StableDiffusionXLTextToImage.from_weights( + "zeromodels/stable-diffusion-xl-base-1.0", load_dtype="float32" +) +``` + +### Verified against diffusers + +The components of `stable-diffusion-xl-base-1.0` and `stable-diffusion-xl-refiner-1.0` +were checked one by one against diffusers 0.39 / transformers in float32 on the same +inputs (the converted float16 weights computed in float32): + +``` +base text_encoder penultimate max|d|=6.1e-05 (values up to 854) +base text_encoder_2 penultimate max|d|=3.1e-05 +base text_encoder_2 pooled max|d|=3.6e-06 +base unet noise_pred (t=981) max|d|=9.3e-06 +base vae decode (512px) max|d|=3.5e-05 +refiner text_encoder_2 penultimate max|d|=3.1e-05 +refiner unet noise_pred (t=181) max|d|=1.2e-05 +``` + +The pipelines (Euler `leading` spacing, classifier-free guidance with the zeroed or +encoded negative branch, the micro-conditioning values, negative prompts, batches, +`strength`, `denoising_end` / `denoising_start`) were checked against +`StableDiffusionXLPipeline` and `StableDiffusionXLImg2ImgPipeline` on small +random-weight models, where the generated images are pixel-identical on torch and jax. +At full size in float16 (the default load; diffusers in float16 with its float32 VAE), +the lighthouse above matches diffusers to PSNR 49.7 dB (98.5% of pixels within 2 uint8 +levels), the base + refiner ensemble to PSNR 51.9 dB (98.8% within 2) and SDXL-Turbo +(1 step, 512px) to PSNR 44.9 dB (92% within 2; a single step amplifies float16 rounding). +The lighthouse, zeromodels (left) and diffusers (right): + +zeromodels (left) and diffusers (right) generations of the lighthouse from stable-diffusion-xl-base-1.0 with the same latent + +## Data Format + +Both `channels_last` and `channels_first` are supported, as for +[Stable Diffusion](stable_diffusion.md#data-format); `generate` always returns +`(batch, H, W, 3)` uint8. + +## Memory and speed + +Measured on the torch backend with the default float16 load (weights resident: 6.8 GB for +a base model, 6.3 GB for the refiner, the VAE in float32): + +| Resolution | Denoising step (guided, batch 1) | VAE decode (float32) | +|---|---|---| +| 512px | peak 7.1 GB, 1.2 s/step on an RTX 4060 Laptop | peak 8.5 GB | +| 768px | peak 8.0 GB | peak 10.5 GB | +| 1024px | peak 10.4 GB | about 14 GB | + +Keras' functional executor keeps every layer output of a graph alive until the graph +returns, so the decoder's feature maps at 1024px add several GB on top of the weights: a +16 GB GPU runs 1024px comfortably, a 12 GB GPU 768px; on 8 GB, generate at 512px, or +denoise on the GPU (`output_type="latent"`) and decode the latent on the CPU (about 100 s +at 1024px) with a second, CPU-only process. In float32 the weights alone take 13.9 GB. +The UNet's attention runs the portable `"sdpa"` math by default, which materializes the +4096 x 4096 float32 logits of the 64 x 64 level at 1024px; loading with +`attn_implementation="fused"` (the backend's fused kernel, see +[`fused_attention`](main_classes.md#fused_attention)) cuts the 1024px step from a 10.3 GB +peak to 7.6 GB on the same GPU. On JAX and TensorFlow the denoiser step is compiled once +per run (see [`BaseDiffusion`](main_classes.md#basediffusion)). + +## Loading Fine-tuned Weights + +Any repo laid out like the hosted ones (`zm_config.json` declaring +`StableDiffusionXLModel`, the weights, `tokenizer.json`) loads with +`from_weights("/")`. The `hf:` prefix raises for diffusion models: convert a +diffusers-format SDXL checkpoint once with +`zeromodels/models/stable_diffusion_xl/convert_stable_diffusion_xl_diffusers_to_keras.py` +(`transfer_stable_diffusion_xl(repo)`, `pip install zeromodels[conversion]`) and host the result. diff --git a/tests/base/model_test_registry.py b/tests/base/model_test_registry.py index 03da6faa..c59c7421 100644 --- a/tests/base/model_test_registry.py +++ b/tests/base/model_test_registry.py @@ -4366,6 +4366,237 @@ "expected_output_shape": dict(_sd_outputs), } +# Stable Diffusion 2.x: the same container / task over the SD 2 configuration (per-level +# attention heads, linear token projection, gelu OpenCLIP-style text tower, "!" padding). +_sd2_tiny = dict( + _sd_tiny, unet_num_attention_heads=(2, 4), unet_use_linear_projection=True +) +MODEL_TEST_CONFIGS["StableDiffusion2Model"] = { + "module": "zeromodels.models.stable_diffusion_2", + "model_cls": "StableDiffusion2Model", + "model_type": "diffusion", + "init_kwargs": dict(_sd2_tiny), + "input_factory": "stable_diffusion_input", + "expected_output_shape": dict(_sd_outputs), +} +MODEL_TEST_CONFIGS["StableDiffusion2TextToImage"] = { + "module": "zeromodels.models.stable_diffusion_2", + "model_cls": "StableDiffusion2TextToImage", + "model_type": "diffusion", + "init_kwargs": dict(_sd2_tiny), + "input_factory": "stable_diffusion_input", + "expected_output_shape": dict(_sd_outputs), +} + +# Stable Diffusion XL: three-level UNet (no attention at the first level, stacked +# transformer blocks, text_time micro-conditioning), two text towers whose penultimate +# states concatenate into the 2048-d context (here 32 + 32 = 64), float32 VAE. +_sdxl_tiny = { + "unet_sample_size": 8, + "unet_down_block_types": ("DownBlock2D", "CrossAttnDownBlock2D"), + "unet_up_block_types": ("CrossAttnUpBlock2D", "UpBlock2D"), + "unet_block_out_channels": (32, 64), + "unet_layers_per_block": 1, + "unet_cross_attention_dim": 64, + "unet_num_attention_heads": (2, 4), + "unet_norm_num_groups": 8, + "unet_use_linear_projection": True, + "unet_transformer_layers_per_block": (1, 2), + "unet_addition_embed_type": "text_time", + "unet_addition_time_embed_dim": 8, + "unet_projection_class_embeddings_input_dim": 16 + 6 * 8, + "unet_num_time_ids": 6, + "unet_text_seq_len": 16, + "vae_sample_size": 16, + "vae_block_out_channels": (16, 32), + "vae_layers_per_block": 1, + "vae_norm_num_groups": 8, + "text_hidden_dim": 32, + "text_num_heads": 2, + "text_num_layers": 2, + "max_seq_len": 16, + "vocab_size": 128, + "text_2_hidden_dim": 32, + "text_2_num_heads": 2, + "text_2_num_layers": 2, + "text_2_projection_dim": 16, + "text_2_max_seq_len": 16, + "text_2_vocab_size": 128, +} +_sdxl_outputs = { + "noise_pred": (2, 8, 8, 4), + "moments": (2, 8, 8, 8), + "image": (2, 16, 16, 3), + "prompt_embeds": (2, 16, 64), + "pooled_prompt_embeds": (2, 16), +} +MODEL_TEST_CONFIGS["StableDiffusionXLModel"] = { + "module": "zeromodels.models.stable_diffusion_xl", + "model_cls": "StableDiffusionXLModel", + "model_type": "diffusion", + "init_kwargs": dict(_sdxl_tiny), + "input_factory": "stable_diffusion_xl_input", + "expected_output_shape": dict(_sdxl_outputs), +} +MODEL_TEST_CONFIGS["StableDiffusionXLTextToImage"] = { + "module": "zeromodels.models.stable_diffusion_xl", + "model_cls": "StableDiffusionXLTextToImage", + "model_type": "diffusion", + "init_kwargs": dict(_sdxl_tiny), + "input_factory": "stable_diffusion_xl_input", + "expected_output_shape": dict(_sdxl_outputs), +} + +# SDXL refiner: the second text tower alone (a 32-d context), a four-level UNet with +# attention at the middle levels, five time ids (size, crop, aesthetic score). +_sdxl_refiner_tiny = { + k: v + for k, v in _sdxl_tiny.items() + if not k.startswith("text_") and k not in ("max_seq_len", "vocab_size") +} | { + "unet_down_block_types": ( + "DownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "DownBlock2D", + ), + "unet_up_block_types": ( + "UpBlock2D", + "CrossAttnUpBlock2D", + "CrossAttnUpBlock2D", + "UpBlock2D", + ), + "unet_block_out_channels": (32, 64, 64, 64), + "unet_cross_attention_dim": 32, + "unet_num_attention_heads": (2, 4, 4, 4), + "unet_transformer_layers_per_block": 2, + "unet_projection_class_embeddings_input_dim": 16 + 5 * 8, + "unet_num_time_ids": 5, + "text_2_hidden_dim": 32, + "text_2_num_heads": 2, + "text_2_num_layers": 2, + "text_2_projection_dim": 16, + "text_2_max_seq_len": 16, + "text_2_vocab_size": 128, +} +_sdxl_refiner_outputs = dict(_sdxl_outputs, prompt_embeds=(2, 16, 32)) +MODEL_TEST_CONFIGS["StableDiffusionXLRefinerModel"] = { + "module": "zeromodels.models.stable_diffusion_xl", + "model_cls": "StableDiffusionXLRefinerModel", + "model_type": "diffusion", + "init_kwargs": dict(_sdxl_refiner_tiny), + "input_factory": "stable_diffusion_xl_refiner_input", + "expected_output_shape": dict(_sdxl_refiner_outputs), +} +MODEL_TEST_CONFIGS["StableDiffusionXLRefinerImageToImage"] = { + "module": "zeromodels.models.stable_diffusion_xl", + "model_cls": "StableDiffusionXLRefinerImageToImage", + "model_type": "diffusion", + "init_kwargs": dict(_sdxl_refiner_tiny), + "input_factory": "stable_diffusion_xl_refiner_input", + "expected_output_shape": dict(_sdxl_refiner_outputs), +} + +# Stable Diffusion 3: the MMDiT (joint text/latent transformer with AdaLN-Zero +# modulation), the 16-channel VAE without quant convs, two CLIP towers with +# projections; the T5 tower is external. SD 3.5 adds RMS qk norms and (medium) the +# dual-attention blocks. +_sd3_tiny = { + "transformer_sample_size": 8, + "transformer_patch_size": 2, + "transformer_in_channels": 4, + "transformer_out_channels": 4, + "transformer_num_layers": 2, + "transformer_attention_head_dim": 8, + "transformer_num_attention_heads": 2, + "transformer_joint_attention_dim": 24, + "transformer_caption_projection_dim": 16, + "transformer_pooled_projection_dim": 24, + "transformer_pos_embed_max_size": 12, + "transformer_text_seq_len": 16 + 8, + "vae_sample_size": 16, + "vae_block_out_channels": (16, 32), + "vae_layers_per_block": 1, + "vae_norm_num_groups": 8, + "vae_latent_channels": 4, + "text_hidden_dim": 16, + "text_num_heads": 2, + "text_num_layers": 2, + "text_projection_dim": 16, + "max_seq_len": 16, + "vocab_size": 128, + "text_2_hidden_dim": 8, + "text_2_num_heads": 2, + "text_2_num_layers": 2, + "text_2_projection_dim": 8, + "text_2_max_seq_len": 16, + "text_2_vocab_size": 128, + "max_sequence_length": 8, +} +_sd3_outputs = { + "noise_pred": (2, 8, 8, 4), + "moments": (2, 8, 8, 8), + "image": (2, 16, 16, 3), + "prompt_embeds": (2, 16, 24), + "pooled_prompt_embeds": (2, 24), +} +MODEL_TEST_CONFIGS["StableDiffusion3Model"] = { + "module": "zeromodels.models.stable_diffusion_3", + "model_cls": "StableDiffusion3Model", + "model_type": "diffusion", + "init_kwargs": dict(_sd3_tiny), + "input_factory": "stable_diffusion_3_input", + "expected_output_shape": dict(_sd3_outputs), +} +MODEL_TEST_CONFIGS["StableDiffusion3TextToImage"] = { + "module": "zeromodels.models.stable_diffusion_3", + "model_cls": "StableDiffusion3TextToImage", + "model_type": "diffusion", + "init_kwargs": dict(_sd3_tiny), + "input_factory": "stable_diffusion_3_input", + "expected_output_shape": dict(_sd3_outputs), +} +# the separately hosted third text encoder (T5 v1.1 XXL layout: gated-GELU blocks) +MODEL_TEST_CONFIGS["SD3T5EncoderModel"] = { + "module": "zeromodels.models.stable_diffusion_3", + "model_cls": "SD3T5EncoderModel", + "model_type": "llm", + "init_kwargs": { + "vocab_size": 128, + "embed_dim": 32, + "key_value_dim": 8, + "mlp_dim": 64, + "num_layers": 2, + "num_heads": 4, + "relative_attention_num_buckets": 8, + "relative_attention_max_distance": 16, + }, + "input_factory": "t5_encoder_input", + "input_factory_kwargs": {"seq_len": 16}, + "expected_output_shape": {"last_hidden_state": (2, 16, 32)}, +} +_sd3_5_tiny = dict( + _sd3_tiny, + transformer_qk_norm="rms_norm", + transformer_dual_attention_layers=(0,), +) +MODEL_TEST_CONFIGS["StableDiffusion3_5Model"] = { + "module": "zeromodels.models.stable_diffusion_3_5", + "model_cls": "StableDiffusion3_5Model", + "model_type": "diffusion", + "init_kwargs": dict(_sd3_5_tiny), + "input_factory": "stable_diffusion_3_input", + "expected_output_shape": dict(_sd3_outputs), +} +MODEL_TEST_CONFIGS["StableDiffusion3_5TextToImage"] = { + "module": "zeromodels.models.stable_diffusion_3_5", + "model_cls": "StableDiffusion3_5TextToImage", + "model_type": "diffusion", + "init_kwargs": dict(_sd3_5_tiny), + "input_factory": "stable_diffusion_3_input", + "expected_output_shape": dict(_sd3_outputs), +} + def get_all_model_ids(): return list(MODEL_TEST_CONFIGS.keys()) diff --git a/tests/fixtures/cross_backend_parity.json b/tests/fixtures/cross_backend_parity.json index 71bd021e..70e27c96 100644 --- a/tests/fixtures/cross_backend_parity.json +++ b/tests/fixtures/cross_backend_parity.json @@ -21755,6 +21755,81 @@ ] } ], + "SD3T5EncoderModel": [ + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + 0.00707, + 0.006957, + 0.009331, + 0.004635, + -0.015012, + -0.007533, + -0.008983, + 0.000647, + 0.005648, + -0.0074, + 0.003372, + -0.011393, + 0.001414, + 0.027524, + -0.005903, + -0.011301, + 0.012131, + 0.013565, + 0.020675, + 0.002784, + 0.028628, + -0.004386, + 0.000328, + -0.004168, + 0.012268, + -0.01554, + -2.8e-05, + -0.012873, + 2.4e-05, + 0.054132, + -0.010922, + 0.005437, + 0.015127, + -0.027579, + -0.002549, + 0.003275, + 0.001096, + 0.010374, + -0.008043, + -0.030602, + 0.000448, + -0.032054, + 0.014059, + -0.003486, + -0.030945, + -0.000766, + -0.015376, + -0.001499, + -0.00409, + 0.001824, + 0.002497, + 0.015728, + 0.018065, + -0.01547, + -0.037632, + -0.002506, + -0.001476, + -0.004758, + 0.0006, + -0.01241, + 0.000607, + -0.003328, + 5.1e-05, + -0.014076 + ] + } + ], "SENetImageClassify": [ { "shape": [ @@ -22605,6 +22680,3954 @@ ] } ], + "StableDiffusion2Model": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015141, + 0.014566, + 0.01507, + -0.013042, + -0.014793, + -0.013933, + 0.01303, + 0.014152, + 0.013037, + 0.014178, + 0.014338, + -0.013882, + -0.013874, + -0.013862, + 0.01476, + 0.014129, + 0.014743, + 0.014246, + 0.014257, + 0.014225, + -0.014007, + -0.013831, + 0.014013, + 0.014092, + 0.014013, + 0.014246, + 0.014301, + 0.014293, + -0.013722, + -0.013693, + -0.013301, + 0.014484, + 0.010029, + 0.014481, + 0.015246, + 0.015245, + -0.013927, + -0.013932, + -0.01385, + 0.014118, + 0.014105, + 0.014152, + 0.014243, + 0.01403, + -0.013841, + -0.013774, + -0.013812, + 0.014105, + 0.014097, + 0.014094, + 0.014242, + 0.014448, + 0.014266, + -0.014007, + -0.013837, + 0.014, + 0.014121, + 0.014003, + 0.014259, + 0.011565, + 0.014413, + -0.011314, + -0.01329, + 0.013126 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051341, + 0.051341, + 0.05132, + 0.051681, + -0.006162, + -0.006151, + -0.006155, + -0.006046, + 0.006829, + 0.006821, + 0.006841, + 0.006748, + 0.010876, + 0.010889, + 0.010892, + 0.010883, + -0.021278, + -0.02128, + -0.021273, + -0.021104, + -0.005376, + -0.005379, + -0.005368, + -0.005528, + -0.005407, + -0.000975, + -0.000971, + -0.00076, + -0.000734, + -0.043369, + -0.043349, + -0.043402, + -0.043551, + 0.051337, + 0.051297, + 0.05151, + 0.051494, + -0.006147, + -0.006312, + -0.006151, + -0.006165, + 0.006816, + 0.006766, + 0.006782, + 0.006824, + 0.006837, + 0.010968, + 0.010946, + 0.0109, + 0.010906, + -0.021234, + -0.021271, + -0.021282, + -0.021255, + -0.005589, + -0.005439, + -0.005423, + -0.005422, + -0.000794, + -0.00072, + -0.000731, + -0.000732, + -0.043442 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.006219, + 0.002704, + 0.003533, + 0.003418, + 0.008001, + 0.005587, + 0.005602, + 0.005686, + 0.008691, + 0.002119, + -0.000714, + 0.001933, + -0.002422, + 0.002237, + 0.001836, + 0.002578, + -0.001956, + 0.00071, + -0.003084, + -0.002698, + -0.005522, + -0.001268, + -0.001686, + -1.2e-05, + -0.004803, + 0.000822, + 0.000538, + 0.013731, + 0.01433, + 0.014584, + 0.015002, + 0.015691, + 0.013323, + 0.015859, + 0.016175, + 0.015331, + 0.004414, + 0.004255, + 0.006079, + 0.002555, + 0.005861, + 0.004837, + 0.005748, + 0.002341, + 0.005368, + -9.1e-05, + 0.001814, + 0.003786, + 0.001644, + 0.001444, + 0.00205, + 0.00168, + 8.1e-05, + 0.000167, + -0.001057, + -0.002688, + -0.000987, + -0.000541, + -0.000336, + -0.002546, + -0.005307, + -0.003004, + -0.006414, + 0.018894 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.012036, + 0.027224, + -0.017103, + 0.027472, + -0.020236, + 0.000332, + -0.016442, + 0.015828, + -0.018295, + 0.014603, + -0.005828, + 0.016616, + -0.00709, + -0.005843, + 0.009211, + 0.000567, + 0.049181, + -0.012505, + 0.000328, + -0.009735, + -0.003474, + -0.007774, + -0.024105, + -0.00845, + -0.020531, + -0.008289, + -0.052391, + 0.003785, + -0.052397, + 0.003749, + -0.019434, + 0.013164, + -0.022706, + 0.014415, + 0.028551, + -0.025012, + 0.017438, + -0.033488, + 0.051644, + -0.027285, + 0.047169, + -0.013412, + 0.049211, + 0.009689, + 0.039967, + 0.02768, + 0.015537, + -0.001814, + -0.01621, + -0.001629, + -0.021368, + 0.015966, + -0.035636, + 0.013777, + -0.038478, + 0.024581, + 0.014664, + 0.015239, + 0.014259, + -0.002774, + 0.006665, + -0.001414, + -0.025644, + 0.009934 + ] + } + ], + "StableDiffusion2TextToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015141, + 0.014566, + 0.01507, + -0.013042, + -0.014793, + -0.013933, + 0.01303, + 0.014152, + 0.013037, + 0.014178, + 0.014338, + -0.013882, + -0.013874, + -0.013862, + 0.01476, + 0.014129, + 0.014743, + 0.014246, + 0.014257, + 0.014225, + -0.014007, + -0.013831, + 0.014013, + 0.014092, + 0.014013, + 0.014246, + 0.014301, + 0.014293, + -0.013722, + -0.013693, + -0.013301, + 0.014484, + 0.010029, + 0.014481, + 0.015246, + 0.015245, + -0.013927, + -0.013932, + -0.01385, + 0.014118, + 0.014105, + 0.014152, + 0.014243, + 0.01403, + -0.013841, + -0.013774, + -0.013812, + 0.014105, + 0.014097, + 0.014094, + 0.014242, + 0.014448, + 0.014266, + -0.014007, + -0.013837, + 0.014, + 0.014121, + 0.014003, + 0.014259, + 0.011565, + 0.014413, + -0.011314, + -0.01329, + 0.013126 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051341, + 0.051341, + 0.05132, + 0.051681, + -0.006162, + -0.006151, + -0.006155, + -0.006046, + 0.006829, + 0.006821, + 0.006841, + 0.006748, + 0.010876, + 0.010889, + 0.010892, + 0.010883, + -0.021278, + -0.02128, + -0.021273, + -0.021104, + -0.005376, + -0.005379, + -0.005368, + -0.005528, + -0.005407, + -0.000975, + -0.000971, + -0.00076, + -0.000734, + -0.043369, + -0.043349, + -0.043402, + -0.043551, + 0.051337, + 0.051297, + 0.05151, + 0.051494, + -0.006147, + -0.006312, + -0.006151, + -0.006165, + 0.006816, + 0.006766, + 0.006782, + 0.006824, + 0.006837, + 0.010968, + 0.010946, + 0.0109, + 0.010906, + -0.021234, + -0.021271, + -0.021282, + -0.021255, + -0.005589, + -0.005439, + -0.005423, + -0.005422, + -0.000794, + -0.00072, + -0.000731, + -0.000732, + -0.043442 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.006219, + 0.002704, + 0.003533, + 0.003418, + 0.008001, + 0.005587, + 0.005602, + 0.005686, + 0.008691, + 0.002119, + -0.000714, + 0.001933, + -0.002422, + 0.002237, + 0.001836, + 0.002578, + -0.001956, + 0.00071, + -0.003084, + -0.002698, + -0.005522, + -0.001268, + -0.001686, + -1.2e-05, + -0.004803, + 0.000822, + 0.000538, + 0.013731, + 0.01433, + 0.014584, + 0.015002, + 0.015691, + 0.013323, + 0.015859, + 0.016175, + 0.015331, + 0.004414, + 0.004255, + 0.006079, + 0.002555, + 0.005861, + 0.004837, + 0.005748, + 0.002341, + 0.005368, + -9.1e-05, + 0.001814, + 0.003786, + 0.001644, + 0.001444, + 0.00205, + 0.00168, + 8.1e-05, + 0.000167, + -0.001057, + -0.002688, + -0.000987, + -0.000541, + -0.000336, + -0.002546, + -0.005307, + -0.003004, + -0.006414, + 0.018894 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.012036, + 0.027224, + -0.017103, + 0.027472, + -0.020236, + 0.000332, + -0.016442, + 0.015828, + -0.018295, + 0.014603, + -0.005828, + 0.016616, + -0.00709, + -0.005843, + 0.009211, + 0.000567, + 0.049181, + -0.012505, + 0.000328, + -0.009735, + -0.003474, + -0.007774, + -0.024105, + -0.00845, + -0.020531, + -0.008289, + -0.052391, + 0.003785, + -0.052397, + 0.003749, + -0.019434, + 0.013164, + -0.022706, + 0.014415, + 0.028551, + -0.025012, + 0.017438, + -0.033488, + 0.051644, + -0.027285, + 0.047169, + -0.013412, + 0.049211, + 0.009689, + 0.039967, + 0.02768, + 0.015537, + -0.001814, + -0.01621, + -0.001629, + -0.021368, + 0.015966, + -0.035636, + 0.013777, + -0.038478, + 0.024581, + 0.014664, + 0.015239, + 0.014259, + -0.002774, + 0.006665, + -0.001414, + -0.025644, + 0.009934 + ] + } + ], + "StableDiffusion3Model": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.014955, + 0.014535, + 0.015039, + -0.01294, + -0.015078, + -0.01354, + 0.012603, + 0.014234, + 0.013028, + 0.014104, + 0.014657, + -0.013961, + -0.01397, + -0.013705, + 0.014085, + 0.014107, + 0.01454, + 0.014684, + 0.014214, + 0.01473, + -0.014538, + -0.013919, + 0.013891, + 0.013894, + 0.013345, + 0.014624, + 0.015308, + 0.015109, + -0.013916, + -0.013549, + -0.013454, + 0.015281, + 0.010778, + 0.014866, + 0.015276, + 0.014992, + -0.014008, + -0.013972, + -0.013274, + 0.014006, + 0.014443, + 0.014339, + 0.013793, + 0.014472, + -0.013407, + -0.013649, + -0.013639, + 0.014388, + 0.014662, + 0.014543, + 0.015311, + 0.014633, + 0.013674, + -0.01387, + -0.01422, + 0.013761, + 0.014448, + 0.013939, + 0.01379, + 0.011211, + 0.013861, + -0.010878, + -0.012983, + 0.013007 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.017274, + 0.019562, + 0.019478, + 0.019517, + 0.016443, + 0.035483, + 0.035683, + 0.035216, + 0.030619, + -0.011854, + -0.012384, + -0.011986, + -0.008521, + 0.008372, + 0.008306, + 0.008393, + 0.006177, + 0.009252, + 0.009079, + 0.008854, + 0.009121, + -0.050123, + -0.050035, + -0.050299, + -0.042257, + -0.050059, + 0.019482, + 0.019644, + 0.015013, + 0.013524, + 0.016211, + 0.016645, + 0.0115, + 0.018172, + 0.019602, + 0.017579, + 0.015945, + 0.01631, + 0.035517, + 0.037192, + 0.036485, + 0.035644, + -0.012204, + -0.015612, + -0.012235, + -0.012165, + -0.012408, + 0.012303, + 0.009171, + 0.008228, + 0.008402, + 0.00942, + 0.008764, + 0.00924, + 0.009194, + -0.047951, + -0.049781, + -0.049926, + -0.050071, + 0.018187, + 0.013454, + 0.013745, + 0.013502, + 0.018325 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.012091, + 0.097546, + 0.1286, + 0.065201, + -0.030095, + 0.150227, + 0.080477, + -0.046941, + 0.100516, + -0.039675, + -0.109624, + -0.035173, + -0.018982, + 0.017097, + 0.0496, + -0.016586, + -0.028893, + -0.045285, + 0.017648, + -0.007875, + -0.051069, + 0.043383, + 0.113692, + -0.142047, + -0.113984, + -0.090202, + -0.088588, + 0.136368, + 0.105976, + 0.16715, + 0.126743, + 0.109739, + -0.003245, + 0.006413, + 0.211325, + 0.010501, + 0.100809, + -0.026821, + -0.011584, + -0.063026, + 0.151531, + 0.07651, + 0.149342, + 0.198445, + -0.054413, + 0.071169, + 0.004813, + -0.014978, + 0.154012, + 0.049127, + -0.056784, + 0.082541, + 0.070801, + 0.150369, + 0.045853, + -0.040565, + 0.095416, + -0.002005, + 0.055735, + 0.052691, + -0.016378, + 0.060451, + 0.060556, + 0.000441 + ] + }, + { + "shape": [ + 2, + 24 + ], + "sample": [ + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829, + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829 + ] + }, + { + "shape": [ + 2, + 16, + 24 + ], + "sample": [ + 0.01991, + 0.004829, + 0.01725, + 0.038287, + -0.006734, + 0.004269, + 0.045626, + 0.02678, + 0.05008, + 0.019808, + 0.082032, + 0.008582, + 0.014972, + 0.010277, + 0.034913, + 0.022746, + 0.015038, + 0.01132, + -0.021046, + -0.056212, + -0.001057, + -0.067862, + -0.024714, + -0.057371, + -0.035649, + -0.058719, + -0.052335, + -0.043667, + -0.032667, + -0.017269, + 0.023394, + -0.025872, + 0.037565, + -0.039031, + 0.032609, + -0.011299, + 0.029473, + 0.035091, + 0.006203, + 0.019498, + 0.028072, + 0.015388, + -0.055906, + 0.0522, + -0.031006, + 0.028746, + -0.001313, + -0.05912, + -0.015437, + -0.072604, + 0.00309, + -0.052105, + -0.060652, + 0.002501, + -0.005301, + -0.016556, + 0.004076, + -0.030116, + 0.015446, + -0.040764, + 0.025618, + -0.082329, + 0.017716, + 0.019545 + ] + } + ], + "StableDiffusion3TextToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.014955, + 0.014535, + 0.015039, + -0.01294, + -0.015078, + -0.01354, + 0.012603, + 0.014234, + 0.013028, + 0.014104, + 0.014657, + -0.013961, + -0.01397, + -0.013705, + 0.014085, + 0.014107, + 0.01454, + 0.014684, + 0.014214, + 0.01473, + -0.014538, + -0.013919, + 0.013891, + 0.013894, + 0.013345, + 0.014624, + 0.015308, + 0.015109, + -0.013916, + -0.013549, + -0.013454, + 0.015281, + 0.010778, + 0.014866, + 0.015276, + 0.014992, + -0.014008, + -0.013972, + -0.013274, + 0.014006, + 0.014443, + 0.014339, + 0.013793, + 0.014472, + -0.013407, + -0.013649, + -0.013639, + 0.014388, + 0.014662, + 0.014543, + 0.015311, + 0.014633, + 0.013674, + -0.01387, + -0.01422, + 0.013761, + 0.014448, + 0.013939, + 0.01379, + 0.011211, + 0.013861, + -0.010878, + -0.012983, + 0.013007 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.017274, + 0.019562, + 0.019478, + 0.019517, + 0.016443, + 0.035483, + 0.035683, + 0.035216, + 0.030619, + -0.011854, + -0.012384, + -0.011986, + -0.008521, + 0.008372, + 0.008306, + 0.008393, + 0.006177, + 0.009252, + 0.009079, + 0.008854, + 0.009121, + -0.050123, + -0.050035, + -0.050299, + -0.042257, + -0.050059, + 0.019482, + 0.019644, + 0.015013, + 0.013524, + 0.016211, + 0.016645, + 0.0115, + 0.018172, + 0.019602, + 0.017579, + 0.015945, + 0.01631, + 0.035517, + 0.037192, + 0.036485, + 0.035644, + -0.012204, + -0.015612, + -0.012235, + -0.012165, + -0.012408, + 0.012303, + 0.009171, + 0.008228, + 0.008402, + 0.00942, + 0.008764, + 0.00924, + 0.009194, + -0.047951, + -0.049781, + -0.049926, + -0.050071, + 0.018187, + 0.013454, + 0.013745, + 0.013502, + 0.018325 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.012091, + 0.097546, + 0.1286, + 0.065201, + -0.030095, + 0.150227, + 0.080477, + -0.046941, + 0.100516, + -0.039675, + -0.109624, + -0.035173, + -0.018982, + 0.017097, + 0.0496, + -0.016586, + -0.028893, + -0.045285, + 0.017648, + -0.007875, + -0.051069, + 0.043383, + 0.113692, + -0.142047, + -0.113984, + -0.090202, + -0.088588, + 0.136368, + 0.105976, + 0.16715, + 0.126743, + 0.109739, + -0.003245, + 0.006413, + 0.211325, + 0.010501, + 0.100809, + -0.026821, + -0.011584, + -0.063026, + 0.151531, + 0.07651, + 0.149342, + 0.198445, + -0.054413, + 0.071169, + 0.004813, + -0.014978, + 0.154012, + 0.049127, + -0.056784, + 0.082541, + 0.070801, + 0.150369, + 0.045853, + -0.040565, + 0.095416, + -0.002005, + 0.055735, + 0.052691, + -0.016378, + 0.060451, + 0.060556, + 0.000441 + ] + }, + { + "shape": [ + 2, + 24 + ], + "sample": [ + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829, + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829 + ] + }, + { + "shape": [ + 2, + 16, + 24 + ], + "sample": [ + 0.01991, + 0.004829, + 0.01725, + 0.038287, + -0.006734, + 0.004269, + 0.045626, + 0.02678, + 0.05008, + 0.019808, + 0.082032, + 0.008582, + 0.014972, + 0.010277, + 0.034913, + 0.022746, + 0.015038, + 0.01132, + -0.021046, + -0.056212, + -0.001057, + -0.067862, + -0.024714, + -0.057371, + -0.035649, + -0.058719, + -0.052335, + -0.043667, + -0.032667, + -0.017269, + 0.023394, + -0.025872, + 0.037565, + -0.039031, + 0.032609, + -0.011299, + 0.029473, + 0.035091, + 0.006203, + 0.019498, + 0.028072, + 0.015388, + -0.055906, + 0.0522, + -0.031006, + 0.028746, + -0.001313, + -0.05912, + -0.015437, + -0.072604, + 0.00309, + -0.052105, + -0.060652, + 0.002501, + -0.005301, + -0.016556, + 0.004076, + -0.030116, + 0.015446, + -0.040764, + 0.025618, + -0.082329, + 0.017716, + 0.019545 + ] + } + ], + "StableDiffusion3_5Model": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.014955, + 0.014535, + 0.015039, + -0.01294, + -0.015078, + -0.01354, + 0.012603, + 0.014234, + 0.013028, + 0.014104, + 0.014657, + -0.013961, + -0.01397, + -0.013705, + 0.014085, + 0.014107, + 0.01454, + 0.014684, + 0.014214, + 0.01473, + -0.014538, + -0.013919, + 0.013891, + 0.013894, + 0.013345, + 0.014624, + 0.015308, + 0.015109, + -0.013916, + -0.013549, + -0.013454, + 0.015281, + 0.010778, + 0.014866, + 0.015276, + 0.014992, + -0.014008, + -0.013972, + -0.013274, + 0.014006, + 0.014443, + 0.014339, + 0.013793, + 0.014472, + -0.013407, + -0.013649, + -0.013639, + 0.014388, + 0.014662, + 0.014543, + 0.015311, + 0.014633, + 0.013674, + -0.01387, + -0.01422, + 0.013761, + 0.014448, + 0.013939, + 0.01379, + 0.011211, + 0.013861, + -0.010878, + -0.012983, + 0.013007 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.017274, + 0.019562, + 0.019478, + 0.019517, + 0.016443, + 0.035483, + 0.035683, + 0.035216, + 0.030619, + -0.011854, + -0.012384, + -0.011986, + -0.008521, + 0.008372, + 0.008306, + 0.008393, + 0.006177, + 0.009252, + 0.009079, + 0.008854, + 0.009121, + -0.050123, + -0.050035, + -0.050299, + -0.042257, + -0.050059, + 0.019482, + 0.019644, + 0.015013, + 0.013524, + 0.016211, + 0.016645, + 0.0115, + 0.018172, + 0.019602, + 0.017579, + 0.015945, + 0.01631, + 0.035517, + 0.037192, + 0.036485, + 0.035644, + -0.012204, + -0.015612, + -0.012235, + -0.012165, + -0.012408, + 0.012303, + 0.009171, + 0.008228, + 0.008402, + 0.00942, + 0.008764, + 0.00924, + 0.009194, + -0.047951, + -0.049781, + -0.049926, + -0.050071, + 0.018187, + 0.013454, + 0.013745, + 0.013502, + 0.018325 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.01323, + 0.098254, + 0.129669, + 0.066239, + -0.028639, + 0.150828, + 0.081944, + -0.045938, + 0.101357, + -0.039674, + -0.10952, + -0.035309, + -0.019402, + 0.016394, + 0.048952, + -0.017228, + -0.029025, + -0.045206, + 0.018156, + -0.007403, + -0.051539, + 0.042579, + 0.113596, + -0.142042, + -0.113807, + -0.089958, + -0.088259, + 0.136318, + 0.105069, + 0.166323, + 0.125724, + 0.108803, + -0.003717, + 0.005988, + 0.211244, + 0.010028, + 0.100874, + -0.026819, + -0.011773, + -0.063357, + 0.151896, + 0.077693, + 0.151372, + 0.199062, + -0.054222, + 0.072171, + 0.006164, + -0.014138, + 0.153723, + 0.049434, + -0.056111, + 0.083136, + 0.071386, + 0.150738, + 0.046609, + -0.039163, + 0.09659, + -0.00176, + 0.055724, + 0.053063, + -0.015067, + 0.061407, + 0.061014, + 0.002211 + ] + }, + { + "shape": [ + 2, + 24 + ], + "sample": [ + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829, + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829 + ] + }, + { + "shape": [ + 2, + 16, + 24 + ], + "sample": [ + 0.01991, + 0.004829, + 0.01725, + 0.038287, + -0.006734, + 0.004269, + 0.045626, + 0.02678, + 0.05008, + 0.019808, + 0.082032, + 0.008582, + 0.014972, + 0.010277, + 0.034913, + 0.022746, + 0.015038, + 0.01132, + -0.021046, + -0.056212, + -0.001057, + -0.067862, + -0.024714, + -0.057371, + -0.035649, + -0.058719, + -0.052335, + -0.043667, + -0.032667, + -0.017269, + 0.023394, + -0.025872, + 0.037565, + -0.039031, + 0.032609, + -0.011299, + 0.029473, + 0.035091, + 0.006203, + 0.019498, + 0.028072, + 0.015388, + -0.055906, + 0.0522, + -0.031006, + 0.028746, + -0.001313, + -0.05912, + -0.015437, + -0.072604, + 0.00309, + -0.052105, + -0.060652, + 0.002501, + -0.005301, + -0.016556, + 0.004076, + -0.030116, + 0.015446, + -0.040764, + 0.025618, + -0.082329, + 0.017716, + 0.019545 + ] + } + ], + "StableDiffusion3_5TextToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.014955, + 0.014535, + 0.015039, + -0.01294, + -0.015078, + -0.01354, + 0.012603, + 0.014234, + 0.013028, + 0.014104, + 0.014657, + -0.013961, + -0.01397, + -0.013705, + 0.014085, + 0.014107, + 0.01454, + 0.014684, + 0.014214, + 0.01473, + -0.014538, + -0.013919, + 0.013891, + 0.013894, + 0.013345, + 0.014624, + 0.015308, + 0.015109, + -0.013916, + -0.013549, + -0.013454, + 0.015281, + 0.010778, + 0.014866, + 0.015276, + 0.014992, + -0.014008, + -0.013972, + -0.013274, + 0.014006, + 0.014443, + 0.014339, + 0.013793, + 0.014472, + -0.013407, + -0.013649, + -0.013639, + 0.014388, + 0.014662, + 0.014543, + 0.015311, + 0.014633, + 0.013674, + -0.01387, + -0.01422, + 0.013761, + 0.014448, + 0.013939, + 0.01379, + 0.011211, + 0.013861, + -0.010878, + -0.012983, + 0.013007 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.017274, + 0.019562, + 0.019478, + 0.019517, + 0.016443, + 0.035483, + 0.035683, + 0.035216, + 0.030619, + -0.011854, + -0.012384, + -0.011986, + -0.008521, + 0.008372, + 0.008306, + 0.008393, + 0.006177, + 0.009252, + 0.009079, + 0.008854, + 0.009121, + -0.050123, + -0.050035, + -0.050299, + -0.042257, + -0.050059, + 0.019482, + 0.019644, + 0.015013, + 0.013524, + 0.016211, + 0.016645, + 0.0115, + 0.018172, + 0.019602, + 0.017579, + 0.015945, + 0.01631, + 0.035517, + 0.037192, + 0.036485, + 0.035644, + -0.012204, + -0.015612, + -0.012235, + -0.012165, + -0.012408, + 0.012303, + 0.009171, + 0.008228, + 0.008402, + 0.00942, + 0.008764, + 0.00924, + 0.009194, + -0.047951, + -0.049781, + -0.049926, + -0.050071, + 0.018187, + 0.013454, + 0.013745, + 0.013502, + 0.018325 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.01323, + 0.098254, + 0.129669, + 0.066239, + -0.028639, + 0.150828, + 0.081944, + -0.045938, + 0.101357, + -0.039674, + -0.10952, + -0.035309, + -0.019402, + 0.016394, + 0.048952, + -0.017228, + -0.029025, + -0.045206, + 0.018156, + -0.007403, + -0.051539, + 0.042579, + 0.113596, + -0.142042, + -0.113807, + -0.089958, + -0.088259, + 0.136318, + 0.105069, + 0.166323, + 0.125724, + 0.108803, + -0.003717, + 0.005988, + 0.211244, + 0.010028, + 0.100874, + -0.026819, + -0.011773, + -0.063357, + 0.151896, + 0.077693, + 0.151372, + 0.199062, + -0.054222, + 0.072171, + 0.006164, + -0.014138, + 0.153723, + 0.049434, + -0.056111, + 0.083136, + 0.071386, + 0.150738, + 0.046609, + -0.039163, + 0.09659, + -0.00176, + 0.055724, + 0.053063, + -0.015067, + 0.061407, + 0.061014, + 0.002211 + ] + }, + { + "shape": [ + 2, + 24 + ], + "sample": [ + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829, + 0.000567, + -0.002695, + -0.000344, + 0.001872, + -0.000507, + -0.002267, + -0.003681, + 0.002676, + 0.004716, + -0.00181, + 0.001028, + 0.001284, + -0.000525, + -0.000201, + 0.001415, + 0.002244, + -0.000562, + 0.001038, + 0.00074, + 0.002248, + 0.000935, + 0.002127, + -0.000101, + 0.000829 + ] + }, + { + "shape": [ + 2, + 16, + 24 + ], + "sample": [ + 0.01991, + 0.004829, + 0.01725, + 0.038287, + -0.006734, + 0.004269, + 0.045626, + 0.02678, + 0.05008, + 0.019808, + 0.082032, + 0.008582, + 0.014972, + 0.010277, + 0.034913, + 0.022746, + 0.015038, + 0.01132, + -0.021046, + -0.056212, + -0.001057, + -0.067862, + -0.024714, + -0.057371, + -0.035649, + -0.058719, + -0.052335, + -0.043667, + -0.032667, + -0.017269, + 0.023394, + -0.025872, + 0.037565, + -0.039031, + 0.032609, + -0.011299, + 0.029473, + 0.035091, + 0.006203, + 0.019498, + 0.028072, + 0.015388, + -0.055906, + 0.0522, + -0.031006, + 0.028746, + -0.001313, + -0.05912, + -0.015437, + -0.072604, + 0.00309, + -0.052105, + -0.060652, + 0.002501, + -0.005301, + -0.016556, + 0.004076, + -0.030116, + 0.015446, + -0.040764, + 0.025618, + -0.082329, + 0.017716, + 0.019545 + ] + } + ], + "StableDiffusionModel": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015141, + 0.014566, + 0.01507, + -0.013042, + -0.014793, + -0.013933, + 0.01303, + 0.014152, + 0.013037, + 0.014178, + 0.014338, + -0.013882, + -0.013874, + -0.013862, + 0.01476, + 0.014129, + 0.014743, + 0.014246, + 0.014257, + 0.014225, + -0.014007, + -0.013831, + 0.014013, + 0.014092, + 0.014013, + 0.014246, + 0.014301, + 0.014293, + -0.013722, + -0.013693, + -0.013301, + 0.014484, + 0.010029, + 0.014481, + 0.015246, + 0.015245, + -0.013927, + -0.013932, + -0.01385, + 0.014118, + 0.014105, + 0.014152, + 0.014243, + 0.01403, + -0.013841, + -0.013774, + -0.013812, + 0.014105, + 0.014097, + 0.014094, + 0.014242, + 0.014448, + 0.014266, + -0.014007, + -0.013837, + 0.014, + 0.014121, + 0.014003, + 0.014259, + 0.011565, + 0.014413, + -0.011314, + -0.01329, + 0.013126 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051341, + 0.051341, + 0.05132, + 0.051681, + -0.006162, + -0.006151, + -0.006155, + -0.006046, + 0.006829, + 0.006821, + 0.006841, + 0.006748, + 0.010876, + 0.010889, + 0.010892, + 0.010883, + -0.021278, + -0.02128, + -0.021273, + -0.021104, + -0.005376, + -0.005379, + -0.005368, + -0.005528, + -0.005407, + -0.000975, + -0.000971, + -0.00076, + -0.000734, + -0.043369, + -0.043349, + -0.043402, + -0.043551, + 0.051337, + 0.051297, + 0.05151, + 0.051494, + -0.006147, + -0.006312, + -0.006151, + -0.006165, + 0.006816, + 0.006766, + 0.006782, + 0.006824, + 0.006837, + 0.010968, + 0.010946, + 0.0109, + 0.010906, + -0.021234, + -0.021271, + -0.021282, + -0.021255, + -0.005589, + -0.005439, + -0.005423, + -0.005422, + -0.000794, + -0.00072, + -0.000731, + -0.000732, + -0.043442 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.006219, + 0.002704, + 0.003533, + 0.003418, + 0.008001, + 0.005587, + 0.005602, + 0.005686, + 0.008691, + 0.002119, + -0.000714, + 0.001933, + -0.002422, + 0.002237, + 0.001836, + 0.002578, + -0.001956, + 0.00071, + -0.003084, + -0.002698, + -0.005522, + -0.001268, + -0.001686, + -1.2e-05, + -0.004803, + 0.000822, + 0.000538, + 0.013731, + 0.01433, + 0.014584, + 0.015002, + 0.015691, + 0.013323, + 0.015859, + 0.016175, + 0.015331, + 0.004414, + 0.004255, + 0.006079, + 0.002555, + 0.005861, + 0.004837, + 0.005748, + 0.002341, + 0.005368, + -9.1e-05, + 0.001814, + 0.003786, + 0.001644, + 0.001444, + 0.00205, + 0.00168, + 8.1e-05, + 0.000167, + -0.001057, + -0.002688, + -0.000987, + -0.000541, + -0.000336, + -0.002546, + -0.005307, + -0.003004, + -0.006414, + 0.018894 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.012037, + 0.027224, + -0.017103, + 0.027472, + -0.020236, + 0.000331, + -0.01644, + 0.015825, + -0.018293, + 0.014603, + -0.005827, + 0.016616, + -0.007089, + -0.005842, + 0.009213, + 0.000568, + 0.049181, + -0.012505, + 0.000328, + -0.009735, + -0.003473, + -0.007774, + -0.024105, + -0.00845, + -0.020531, + -0.008289, + -0.052391, + 0.003785, + -0.052397, + 0.003749, + -0.019436, + 0.013164, + -0.022708, + 0.014415, + 0.028549, + -0.025015, + 0.017437, + -0.03349, + 0.051644, + -0.027284, + 0.04717, + -0.013412, + 0.04921, + 0.009694, + 0.039967, + 0.027683, + 0.015538, + -0.001814, + -0.016217, + -0.001629, + -0.021373, + 0.015966, + -0.035635, + 0.013777, + -0.038477, + 0.024585, + 0.014664, + 0.015243, + 0.014259, + -0.002774, + 0.006661, + -0.001414, + -0.025646, + 0.00993 + ] + } + ], + "StableDiffusionTextToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015141, + 0.014566, + 0.01507, + -0.013042, + -0.014793, + -0.013933, + 0.01303, + 0.014152, + 0.013037, + 0.014178, + 0.014338, + -0.013882, + -0.013874, + -0.013862, + 0.01476, + 0.014129, + 0.014743, + 0.014246, + 0.014257, + 0.014225, + -0.014007, + -0.013831, + 0.014013, + 0.014092, + 0.014013, + 0.014246, + 0.014301, + 0.014293, + -0.013722, + -0.013693, + -0.013301, + 0.014484, + 0.010029, + 0.014481, + 0.015246, + 0.015245, + -0.013927, + -0.013932, + -0.01385, + 0.014118, + 0.014105, + 0.014152, + 0.014243, + 0.01403, + -0.013841, + -0.013774, + -0.013812, + 0.014105, + 0.014097, + 0.014094, + 0.014242, + 0.014448, + 0.014266, + -0.014007, + -0.013837, + 0.014, + 0.014121, + 0.014003, + 0.014259, + 0.011565, + 0.014413, + -0.011314, + -0.01329, + 0.013126 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051341, + 0.051341, + 0.05132, + 0.051681, + -0.006162, + -0.006151, + -0.006155, + -0.006046, + 0.006829, + 0.006821, + 0.006841, + 0.006748, + 0.010876, + 0.010889, + 0.010892, + 0.010883, + -0.021278, + -0.02128, + -0.021273, + -0.021104, + -0.005376, + -0.005379, + -0.005368, + -0.005528, + -0.005407, + -0.000975, + -0.000971, + -0.00076, + -0.000734, + -0.043369, + -0.043349, + -0.043402, + -0.043551, + 0.051337, + 0.051297, + 0.05151, + 0.051494, + -0.006147, + -0.006312, + -0.006151, + -0.006165, + 0.006816, + 0.006766, + 0.006782, + 0.006824, + 0.006837, + 0.010968, + 0.010946, + 0.0109, + 0.010906, + -0.021234, + -0.021271, + -0.021282, + -0.021255, + -0.005589, + -0.005439, + -0.005423, + -0.005422, + -0.000794, + -0.00072, + -0.000731, + -0.000732, + -0.043442 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.006219, + 0.002704, + 0.003533, + 0.003418, + 0.008001, + 0.005587, + 0.005602, + 0.005686, + 0.008691, + 0.002119, + -0.000714, + 0.001933, + -0.002422, + 0.002237, + 0.001836, + 0.002578, + -0.001956, + 0.00071, + -0.003084, + -0.002698, + -0.005522, + -0.001268, + -0.001686, + -1.2e-05, + -0.004803, + 0.000822, + 0.000538, + 0.013731, + 0.01433, + 0.014584, + 0.015002, + 0.015691, + 0.013323, + 0.015859, + 0.016175, + 0.015331, + 0.004414, + 0.004255, + 0.006079, + 0.002555, + 0.005861, + 0.004837, + 0.005748, + 0.002341, + 0.005368, + -9.1e-05, + 0.001814, + 0.003786, + 0.001644, + 0.001444, + 0.00205, + 0.00168, + 8.1e-05, + 0.000167, + -0.001057, + -0.002688, + -0.000987, + -0.000541, + -0.000336, + -0.002546, + -0.005307, + -0.003004, + -0.006414, + 0.018894 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.012037, + 0.027224, + -0.017103, + 0.027472, + -0.020236, + 0.000331, + -0.01644, + 0.015825, + -0.018293, + 0.014603, + -0.005827, + 0.016616, + -0.007089, + -0.005842, + 0.009213, + 0.000568, + 0.049181, + -0.012505, + 0.000328, + -0.009735, + -0.003473, + -0.007774, + -0.024105, + -0.00845, + -0.020531, + -0.008289, + -0.052391, + 0.003785, + -0.052397, + 0.003749, + -0.019436, + 0.013164, + -0.022708, + 0.014415, + 0.028549, + -0.025015, + 0.017437, + -0.03349, + 0.051644, + -0.027284, + 0.04717, + -0.013412, + 0.04921, + 0.009694, + 0.039967, + 0.027683, + 0.015538, + -0.001814, + -0.016217, + -0.001629, + -0.021373, + 0.015966, + -0.035635, + 0.013777, + -0.038477, + 0.024585, + 0.014664, + 0.015243, + 0.014259, + -0.002774, + 0.006661, + -0.001414, + -0.025646, + 0.00993 + ] + } + ], + "StableDiffusionXLModel": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015117, + 0.014538, + 0.015099, + -0.013055, + -0.014773, + -0.013969, + 0.013046, + 0.014114, + 0.013047, + 0.014206, + 0.01431, + -0.013879, + -0.013879, + -0.013813, + 0.014727, + 0.014137, + 0.014728, + 0.014265, + 0.014223, + 0.014277, + -0.013994, + -0.013862, + 0.013943, + 0.014071, + 0.01396, + 0.0142, + 0.014354, + 0.014264, + -0.01367, + -0.013662, + -0.013273, + 0.014495, + 0.009997, + 0.014473, + 0.015255, + 0.015218, + -0.013923, + -0.01392, + -0.013851, + 0.014064, + 0.014162, + 0.014085, + 0.014243, + 0.014099, + -0.013891, + -0.013708, + -0.013863, + 0.014129, + 0.01412, + 0.014114, + 0.014283, + 0.014466, + 0.014317, + -0.014031, + -0.013894, + 0.013986, + 0.01413, + 0.014016, + 0.014216, + 0.011546, + 0.014328, + -0.011294, + -0.013285, + 0.013122 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051605, + 0.051336, + 0.051333, + 0.051329, + 0.051681, + -0.006148, + -0.006164, + -0.006143, + -0.006038, + 0.006828, + 0.006835, + 0.006854, + 0.006746, + 0.01086, + 0.010873, + 0.010902, + 0.010893, + -0.021265, + -0.021276, + -0.021282, + -0.021111, + -0.005397, + -0.005385, + -0.005377, + -0.005534, + -0.005419, + -0.000964, + -0.000972, + -0.000765, + -0.000732, + -0.043374, + -0.043356, + -0.043409, + -0.043543, + 0.051326, + 0.051296, + 0.051495, + 0.051486, + -0.006144, + -0.006299, + -0.006143, + -0.006163, + 0.006826, + 0.006775, + 0.006793, + 0.006821, + 0.006818, + 0.010973, + 0.010951, + 0.010897, + 0.010886, + -0.02123, + -0.021273, + -0.021282, + -0.021264, + -0.005587, + -0.005444, + -0.005419, + -0.005421, + -0.000796, + -0.000722, + -0.000731, + -0.00073, + -0.043448 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.005391, + 0.001505, + 0.002643, + 0.002754, + 0.005352, + 0.003054, + 0.002996, + 0.002889, + 0.006612, + 0.001375, + -0.00157, + 0.001698, + -0.003237, + 0.002067, + 0.001719, + 0.002174, + -0.003029, + -0.000492, + -0.005456, + -0.005545, + -0.007489, + -0.002657, + -0.003191, + -0.000723, + -0.006884, + -0.000489, + -0.000329, + 0.013296, + 0.014399, + 0.012816, + 0.013579, + 0.014752, + 0.014613, + 0.016026, + 0.016497, + 0.015358, + 0.001693, + 0.000713, + 0.00396, + 0.001166, + 0.003711, + 0.002246, + 0.002561, + 0.000427, + 0.00251, + -0.001125, + 0.002253, + 0.004757, + 0.001045, + 0.000851, + 0.001967, + 0.001514, + -0.000799, + -5.2e-05, + -0.002928, + -0.003612, + -0.003489, + -0.002121, + -0.001713, + -0.003564, + -0.005562, + -0.003289, + -0.007317, + 0.016749 + ] + }, + { + "shape": [ + 2, + 16 + ], + "sample": [ + -0.001651, + 0.004384, + 0.000777, + 0.001588, + 0.001558, + 0.002165, + -0.002537, + 0.002124, + -0.002929, + 0.004796, + 0.002518, + -0.000757, + -0.006795, + -0.003231, + -0.005023, + 0.002541, + -0.001651, + 0.004384, + 0.000777, + 0.001588, + 0.001558, + 0.002165, + -0.002537, + 0.002124, + -0.002929, + 0.004796, + 0.002518, + -0.000757, + -0.006795, + -0.003231, + -0.005023, + 0.002541 + ] + }, + { + "shape": [ + 2, + 16, + 64 + ], + "sample": [ + 0.022025, + -0.090905, + -0.005076, + 0.007606, + 0.008709, + -0.020836, + -0.006304, + 0.056435, + -0.066661, + -0.024251, + -0.024757, + -0.038454, + 0.083587, + -0.02392, + 0.07096, + -0.037477, + -0.053752, + 0.016964, + -0.025098, + -0.000555, + 0.048412, + -0.056334, + -0.012582, + 0.013893, + -0.012235, + 0.048392, + 0.008929, + -0.023006, + -0.018066, + 0.025023, + -0.021386, + -0.020507, + -0.118211, + -0.040016, + -0.012872, + -0.005505, + 0.041438, + -0.029135, + 0.015047, + 0.003094, + 0.003135, + -0.052091, + -0.026491, + 0.025553, + 0.032734, + -0.047294, + -0.003716, + -0.016993, + 0.002541, + -0.002171, + -0.05959, + -0.001694, + 0.016903, + -0.071585, + 0.046993, + 0.011406, + -0.007823, + -0.061899, + 0.043031, + -0.024788, + -0.030914, + -0.014168, + -0.01589, + -0.052181 + ] + } + ], + "StableDiffusionXLRefinerImageToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015128, + 0.014555, + 0.015074, + -0.013027, + -0.014804, + -0.013913, + 0.013123, + 0.014134, + 0.013115, + 0.014175, + 0.01428, + -0.01386, + -0.013889, + -0.013869, + 0.014721, + 0.014142, + 0.014713, + 0.014238, + 0.014288, + 0.014254, + -0.013996, + -0.013801, + 0.013976, + 0.014124, + 0.013997, + 0.014224, + 0.01436, + 0.014271, + -0.013665, + -0.013693, + -0.013272, + 0.014464, + 0.010032, + 0.014493, + 0.015215, + 0.015217, + -0.013947, + -0.013935, + -0.013853, + 0.014106, + 0.014111, + 0.014094, + 0.014294, + 0.014129, + -0.01385, + -0.013772, + -0.01386, + 0.014143, + 0.014081, + 0.014134, + 0.014289, + 0.014426, + 0.014263, + -0.014039, + -0.013828, + 0.013967, + 0.014145, + 0.013975, + 0.014264, + 0.011534, + 0.014323, + -0.011284, + -0.013253, + 0.013137 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051339, + 0.05134, + 0.051333, + 0.051681, + -0.006151, + -0.00616, + -0.006148, + -0.006045, + 0.006832, + 0.006832, + 0.006855, + 0.00675, + 0.0109, + 0.010893, + 0.010888, + 0.010893, + -0.02127, + -0.021268, + -0.021271, + -0.021117, + -0.005368, + -0.00538, + -0.005377, + -0.005524, + -0.005409, + -0.000983, + -0.000983, + -0.000756, + -0.000733, + -0.043366, + -0.043358, + -0.043399, + -0.043558, + 0.051339, + 0.051289, + 0.051522, + 0.051488, + -0.006163, + -0.006294, + -0.006157, + -0.006156, + 0.006832, + 0.006773, + 0.006796, + 0.006826, + 0.006832, + 0.010985, + 0.010941, + 0.01088, + 0.010879, + -0.021234, + -0.021277, + -0.021264, + -0.02128, + -0.005601, + -0.005452, + -0.005426, + -0.00543, + -0.000791, + -0.00072, + -0.000728, + -0.000727, + -0.043441 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.000481, + -0.002919, + -0.00179, + -0.000417, + 0.003282, + 0.003562, + 0.004124, + 0.003136, + 0.003245, + -0.009335, + -0.007721, + -0.008315, + -0.007362, + -0.010133, + -0.009758, + -0.009103, + -0.008528, + -0.010434, + -0.010003, + -0.007856, + -0.008889, + -0.008763, + -0.01003, + -0.008096, + -0.009952, + -0.006777, + -0.007405, + 0.009624, + 0.010534, + 0.013684, + 0.014548, + 0.014422, + 0.008618, + 0.009694, + 0.008467, + 0.007777, + 0.00285, + 0.003307, + 0.004301, + 0.00407, + 0.003844, + 0.0041, + 0.002652, + 0.003409, + 0.001842, + -0.00776, + -0.008313, + -0.003254, + -0.007042, + -0.009972, + -0.009185, + -0.0035, + -0.008157, + -0.009825, + -0.008416, + -0.006144, + -0.007326, + -0.007724, + -0.008343, + -0.006462, + -0.009628, + -0.010498, + -0.010122, + 0.016935 + ] + }, + { + "shape": [ + 2, + 16 + ], + "sample": [ + -0.004386, + 0.005251, + -0.001644, + 0.002351, + 0.002163, + 0.003551, + -0.002766, + 0.001111, + -0.001758, + 0.005881, + -0.000112, + 0.002642, + -0.006228, + -0.002458, + -0.002431, + -0.000691, + -0.004386, + 0.005251, + -0.001644, + 0.002351, + 0.002163, + 0.003551, + -0.002766, + 0.001111, + -0.001758, + 0.005881, + -0.000112, + 0.002642, + -0.006228, + -0.002458, + -0.002431, + -0.000691 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.07549, + -0.036797, + -0.06729, + -0.010316, + -0.087311, + 0.026806, + 0.046183, + 0.080015, + 0.028766, + 0.010841, + -0.001441, + 0.025617, + -0.022712, + 0.002993, + 0.03557, + 0.032339, + 0.039893, + -0.003567, + -0.040687, + -0.018836, + -0.055347, + -0.010683, + 0.006929, + 0.011367, + 0.034529, + -0.008013, + 0.033182, + -0.019324, + 0.047424, + -0.047148, + -0.093025, + 0.054083, + -0.060365, + 0.062953, + 0.030306, + 0.02033, + 0.052184, + 0.016348, + -0.005382, + -0.010774, + 0.03376, + -0.011427, + -0.024654, + -0.077269, + -0.040009, + -0.079375, + -0.008919, + 0.042451, + -0.07994, + 0.030593, + -0.049702, + -0.078214, + 0.049456, + -0.093707, + 0.04243, + -0.034284, + -0.032833, + -0.028629, + -0.038116, + 0.025034, + 0.0324, + -0.014014, + 0.029427, + -0.024193 + ] + } + ], + "StableDiffusionXLRefinerModel": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015128, + 0.014555, + 0.015074, + -0.013027, + -0.014804, + -0.013913, + 0.013123, + 0.014134, + 0.013115, + 0.014175, + 0.01428, + -0.01386, + -0.013889, + -0.013869, + 0.014721, + 0.014142, + 0.014713, + 0.014238, + 0.014288, + 0.014254, + -0.013996, + -0.013801, + 0.013976, + 0.014124, + 0.013997, + 0.014224, + 0.01436, + 0.014271, + -0.013665, + -0.013693, + -0.013272, + 0.014464, + 0.010032, + 0.014493, + 0.015215, + 0.015217, + -0.013947, + -0.013935, + -0.013853, + 0.014106, + 0.014111, + 0.014094, + 0.014294, + 0.014129, + -0.01385, + -0.013772, + -0.01386, + 0.014143, + 0.014081, + 0.014134, + 0.014289, + 0.014426, + 0.014263, + -0.014039, + -0.013828, + 0.013967, + 0.014145, + 0.013975, + 0.014264, + 0.011534, + 0.014323, + -0.011284, + -0.013253, + 0.013137 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051589, + 0.051339, + 0.05134, + 0.051333, + 0.051681, + -0.006151, + -0.00616, + -0.006148, + -0.006045, + 0.006832, + 0.006832, + 0.006855, + 0.00675, + 0.0109, + 0.010893, + 0.010888, + 0.010893, + -0.02127, + -0.021268, + -0.021271, + -0.021117, + -0.005368, + -0.00538, + -0.005377, + -0.005524, + -0.005409, + -0.000983, + -0.000983, + -0.000756, + -0.000733, + -0.043366, + -0.043358, + -0.043399, + -0.043558, + 0.051339, + 0.051289, + 0.051522, + 0.051488, + -0.006163, + -0.006294, + -0.006157, + -0.006156, + 0.006832, + 0.006773, + 0.006796, + 0.006826, + 0.006832, + 0.010985, + 0.010941, + 0.01088, + 0.010879, + -0.021234, + -0.021277, + -0.021264, + -0.02128, + -0.005601, + -0.005452, + -0.005426, + -0.00543, + -0.000791, + -0.00072, + -0.000728, + -0.000727, + -0.043441 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.000481, + -0.002919, + -0.00179, + -0.000417, + 0.003282, + 0.003562, + 0.004124, + 0.003136, + 0.003245, + -0.009335, + -0.007721, + -0.008315, + -0.007362, + -0.010133, + -0.009758, + -0.009103, + -0.008528, + -0.010434, + -0.010003, + -0.007856, + -0.008889, + -0.008763, + -0.01003, + -0.008096, + -0.009952, + -0.006777, + -0.007405, + 0.009624, + 0.010534, + 0.013684, + 0.014548, + 0.014422, + 0.008618, + 0.009694, + 0.008467, + 0.007777, + 0.00285, + 0.003307, + 0.004301, + 0.00407, + 0.003844, + 0.0041, + 0.002652, + 0.003409, + 0.001842, + -0.00776, + -0.008313, + -0.003254, + -0.007042, + -0.009972, + -0.009185, + -0.0035, + -0.008157, + -0.009825, + -0.008416, + -0.006144, + -0.007326, + -0.007724, + -0.008343, + -0.006462, + -0.009628, + -0.010498, + -0.010122, + 0.016935 + ] + }, + { + "shape": [ + 2, + 16 + ], + "sample": [ + -0.004386, + 0.005251, + -0.001644, + 0.002351, + 0.002163, + 0.003551, + -0.002766, + 0.001111, + -0.001758, + 0.005881, + -0.000112, + 0.002642, + -0.006228, + -0.002458, + -0.002431, + -0.000691, + -0.004386, + 0.005251, + -0.001644, + 0.002351, + 0.002163, + 0.003551, + -0.002766, + 0.001111, + -0.001758, + 0.005881, + -0.000112, + 0.002642, + -0.006228, + -0.002458, + -0.002431, + -0.000691 + ] + }, + { + "shape": [ + 2, + 16, + 32 + ], + "sample": [ + -0.07549, + -0.036797, + -0.06729, + -0.010316, + -0.087311, + 0.026806, + 0.046183, + 0.080015, + 0.028766, + 0.010841, + -0.001441, + 0.025617, + -0.022712, + 0.002993, + 0.03557, + 0.032339, + 0.039893, + -0.003567, + -0.040687, + -0.018836, + -0.055347, + -0.010683, + 0.006929, + 0.011367, + 0.034529, + -0.008013, + 0.033182, + -0.019324, + 0.047424, + -0.047148, + -0.093025, + 0.054083, + -0.060365, + 0.062953, + 0.030306, + 0.02033, + 0.052184, + 0.016348, + -0.005382, + -0.010774, + 0.03376, + -0.011427, + -0.024654, + -0.077269, + -0.040009, + -0.079375, + -0.008919, + 0.042451, + -0.07994, + 0.030593, + -0.049702, + -0.078214, + 0.049456, + -0.093707, + 0.04243, + -0.034284, + -0.032833, + -0.028629, + -0.038116, + 0.025034, + 0.0324, + -0.014014, + 0.029427, + -0.024193 + ] + } + ], + "StableDiffusionXLTextToImage": [ + { + "shape": [ + 2, + 16, + 16, + 3 + ], + "sample": [ + 0.015117, + 0.014538, + 0.015099, + -0.013055, + -0.014773, + -0.013969, + 0.013046, + 0.014114, + 0.013047, + 0.014206, + 0.01431, + -0.013879, + -0.013879, + -0.013813, + 0.014727, + 0.014137, + 0.014728, + 0.014265, + 0.014223, + 0.014277, + -0.013994, + -0.013862, + 0.013943, + 0.014071, + 0.01396, + 0.0142, + 0.014354, + 0.014264, + -0.01367, + -0.013662, + -0.013273, + 0.014495, + 0.009997, + 0.014473, + 0.015255, + 0.015218, + -0.013923, + -0.01392, + -0.013851, + 0.014064, + 0.014162, + 0.014085, + 0.014243, + 0.014099, + -0.013891, + -0.013708, + -0.013863, + 0.014129, + 0.01412, + 0.014114, + 0.014283, + 0.014466, + 0.014317, + -0.014031, + -0.013894, + 0.013986, + 0.01413, + 0.014016, + 0.014216, + 0.011546, + 0.014328, + -0.011294, + -0.013285, + 0.013122 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 8 + ], + "sample": [ + 0.051605, + 0.051336, + 0.051333, + 0.051329, + 0.051681, + -0.006148, + -0.006164, + -0.006143, + -0.006038, + 0.006828, + 0.006835, + 0.006854, + 0.006746, + 0.01086, + 0.010873, + 0.010902, + 0.010893, + -0.021265, + -0.021276, + -0.021282, + -0.021111, + -0.005397, + -0.005385, + -0.005377, + -0.005534, + -0.005419, + -0.000964, + -0.000972, + -0.000765, + -0.000732, + -0.043374, + -0.043356, + -0.043409, + -0.043543, + 0.051326, + 0.051296, + 0.051495, + 0.051486, + -0.006144, + -0.006299, + -0.006143, + -0.006163, + 0.006826, + 0.006775, + 0.006793, + 0.006821, + 0.006818, + 0.010973, + 0.010951, + 0.010897, + 0.010886, + -0.02123, + -0.021273, + -0.021282, + -0.021264, + -0.005587, + -0.005444, + -0.005419, + -0.005421, + -0.000796, + -0.000722, + -0.000731, + -0.00073, + -0.043448 + ] + }, + { + "shape": [ + 2, + 8, + 8, + 4 + ], + "sample": [ + 0.005391, + 0.001505, + 0.002643, + 0.002754, + 0.005352, + 0.003054, + 0.002996, + 0.002889, + 0.006612, + 0.001375, + -0.00157, + 0.001698, + -0.003237, + 0.002067, + 0.001719, + 0.002174, + -0.003029, + -0.000492, + -0.005456, + -0.005545, + -0.007489, + -0.002657, + -0.003191, + -0.000723, + -0.006884, + -0.000489, + -0.000329, + 0.013296, + 0.014399, + 0.012816, + 0.013579, + 0.014752, + 0.014613, + 0.016026, + 0.016497, + 0.015358, + 0.001693, + 0.000713, + 0.00396, + 0.001166, + 0.003711, + 0.002246, + 0.002561, + 0.000427, + 0.00251, + -0.001125, + 0.002253, + 0.004757, + 0.001045, + 0.000851, + 0.001967, + 0.001514, + -0.000799, + -5.2e-05, + -0.002928, + -0.003612, + -0.003489, + -0.002121, + -0.001713, + -0.003564, + -0.005562, + -0.003289, + -0.007317, + 0.016749 + ] + }, + { + "shape": [ + 2, + 16 + ], + "sample": [ + -0.001651, + 0.004384, + 0.000777, + 0.001588, + 0.001558, + 0.002165, + -0.002537, + 0.002124, + -0.002929, + 0.004796, + 0.002518, + -0.000757, + -0.006795, + -0.003231, + -0.005023, + 0.002541, + -0.001651, + 0.004384, + 0.000777, + 0.001588, + 0.001558, + 0.002165, + -0.002537, + 0.002124, + -0.002929, + 0.004796, + 0.002518, + -0.000757, + -0.006795, + -0.003231, + -0.005023, + 0.002541 + ] + }, + { + "shape": [ + 2, + 16, + 64 + ], + "sample": [ + 0.022025, + -0.090905, + -0.005076, + 0.007606, + 0.008709, + -0.020836, + -0.006304, + 0.056435, + -0.066661, + -0.024251, + -0.024757, + -0.038454, + 0.083587, + -0.02392, + 0.07096, + -0.037477, + -0.053752, + 0.016964, + -0.025098, + -0.000555, + 0.048412, + -0.056334, + -0.012582, + 0.013893, + -0.012235, + 0.048392, + 0.008929, + -0.023006, + -0.018066, + 0.025023, + -0.021386, + -0.020507, + -0.118211, + -0.040016, + -0.012872, + -0.005505, + 0.041438, + -0.029135, + 0.015047, + 0.003094, + 0.003135, + -0.052091, + -0.026491, + 0.025553, + 0.032734, + -0.047294, + -0.003716, + -0.016993, + 0.002541, + -0.002171, + -0.05959, + -0.001694, + 0.016903, + -0.071585, + 0.046993, + 0.011406, + -0.007823, + -0.061899, + 0.043031, + -0.024788, + -0.030914, + -0.014168, + -0.01589, + -0.052181 + ] + } + ], "SwinImageClassify": [ { "shape": [ diff --git a/tests/fixtures/dummy_inputs.py b/tests/fixtures/dummy_inputs.py index 85168148..3d926da8 100644 --- a/tests/fixtures/dummy_inputs.py +++ b/tests/fixtures/dummy_inputs.py @@ -160,6 +160,82 @@ def stable_diffusion_input( } +def stable_diffusion_xl_input( + batch_size=2, + sample_size=8, + image_size=16, + max_seq_len=16, + cross_attention_dim=64, + pooled_dim=16, + num_time_ids=6, +): + # SDXL adds the UNet's text_time micro-conditioning (pooled text embedding + + # size / crop ids) and the second text tower's ids to the SD 1.x paths + return { + "sample": ops.ones((batch_size, sample_size, sample_size, 4)), + "timestep": ops.ones((batch_size,)), + "encoder_hidden_states": ops.ones( + (batch_size, max_seq_len, cross_attention_dim) + ), + "text_embeds": ops.ones((batch_size, pooled_dim)), + "time_ids": ops.ones((batch_size, num_time_ids)), + "image": ops.ones((batch_size, image_size, image_size, 3)), + "latent": ops.ones((batch_size, sample_size, sample_size, 4)), + "token_ids": ops.ones((batch_size, max_seq_len), dtype="int32"), + "token_ids_2": ops.ones((batch_size, max_seq_len), dtype="int32"), + "padding_mask": ops.ones((batch_size, max_seq_len), dtype="int32"), + } + + +def stable_diffusion_xl_refiner_input( + batch_size=2, + sample_size=8, + image_size=16, + max_seq_len=16, + cross_attention_dim=32, + pooled_dim=16, + num_time_ids=5, +): + # the refiner has no first text tower: no token_ids, five time ids + inputs = stable_diffusion_xl_input( + batch_size, + sample_size, + image_size, + max_seq_len, + cross_attention_dim, + pooled_dim, + num_time_ids, + ) + inputs.pop("token_ids") + return inputs + + +def stable_diffusion_3_input( + batch_size=2, + sample_size=8, + image_size=16, + max_seq_len=16, + text_seq_len=24, + joint_attention_dim=24, + pooled_dim=24, +): + # the MMDiT takes the text tokens (CLIP + T5 length) at the joint width and + # the pooled CLIP embeddings; the VAE and the two CLIP towers as in SDXL + return { + "sample": ops.ones((batch_size, sample_size, sample_size, 4)), + "timestep": ops.ones((batch_size,)), + "encoder_hidden_states": ops.ones( + (batch_size, text_seq_len, joint_attention_dim) + ), + "pooled_projections": ops.ones((batch_size, pooled_dim)), + "image": ops.ones((batch_size, image_size, image_size, 3)), + "latent": ops.ones((batch_size, sample_size, sample_size, 4)), + "token_ids": ops.ones((batch_size, max_seq_len), dtype="int32"), + "token_ids_2": ops.ones((batch_size, max_seq_len), dtype="int32"), + "padding_mask": ops.ones((batch_size, max_seq_len), dtype="int32"), + } + + def tips_v2_text_input(batch_size=2, max_seq_len=16): return { "token_ids": ops.ones((batch_size, max_seq_len), dtype="int32"), diff --git a/tests/integration/test_auto_registry.py b/tests/integration/test_auto_registry.py index cd446bfb..63e3469f 100644 --- a/tests/integration/test_auto_registry.py +++ b/tests/integration/test_auto_registry.py @@ -42,6 +42,7 @@ "PanopticSegment", "UniversalSegment", "TextToImage", + "ImageToImage", "DptDepthEstimation", "DptSemanticSegment", "DptDensePredict", @@ -276,10 +277,11 @@ def test_every_table_value_resolves_and_matches_its_task(): # Redundant transformers-named alias whose model_type ("grounding-dino") already maps to # the zeromodels-convention sibling GroundingDinoDetect. "GroundingDinoForObjectDetection", - # A component of the Stable Diffusion container (loaded through StableDiffusionModel / - # StableDiffusionTextToImage), never a hosted repo of its own; keeps diffusers' name, - # which carries no task suffix. + # Components of the Stable Diffusion containers (loaded through the family's + # XModel / XTextToImage), never hosted repos of their own; keep diffusers' names, + # which carry no task suffix. "AutoencoderKL", + "SD3Transformer2DModel", } diff --git a/website/mkdocs.yml b/website/mkdocs.yml index 88c7fdab..4803be20 100644 --- a/website/mkdocs.yml +++ b/website/mkdocs.yml @@ -231,6 +231,10 @@ nav: - SigLIP: siglip.md - SigLIP 2: siglip2.md - Stable Diffusion: stable_diffusion.md + - Stable Diffusion 2: stable_diffusion_2.md + - Stable Diffusion XL: stable_diffusion_xl.md + - Stable Diffusion 3: stable_diffusion_3.md + - Stable Diffusion 3.5: stable_diffusion_3_5.md - TIPSv2: tipsv2.md - Loading Weights: loading_weights.md diff --git a/zeromodels/auto/auto_mapping_names.py b/zeromodels/auto/auto_mapping_names.py index 41ffd264..514fd155 100644 --- a/zeromodels/auto/auto_mapping_names.py +++ b/zeromodels/auto/auto_mapping_names.py @@ -75,6 +75,7 @@ "table-transformer": "TableTransformerDetect", }, "EncoderModel": { + "stable_diffusion_3_t5_encoder": "SD3T5EncoderModel", "t5": "T5EncoderModel", }, "ImageClassify": { @@ -271,6 +272,11 @@ "siglip": "SigLIPModel", "siglip2": "SigLIP2Model", "stable_diffusion": "StableDiffusionModel", + "stable_diffusion_2": "StableDiffusion2Model", + "stable_diffusion_xl": "StableDiffusionXLModel", + "stable_diffusion_xl_refiner": "StableDiffusionXLRefinerModel", + "stable_diffusion_3": "StableDiffusion3Model", + "stable_diffusion_3_5": "StableDiffusion3_5Model", "swin": "SwinModel", "swinv2": "SwinV2Model", "t5": "T5Model", @@ -399,6 +405,13 @@ }, "TextToImage": { "stable_diffusion": "StableDiffusionTextToImage", + "stable_diffusion_2": "StableDiffusion2TextToImage", + "stable_diffusion_xl": "StableDiffusionXLTextToImage", + "stable_diffusion_3": "StableDiffusion3TextToImage", + "stable_diffusion_3_5": "StableDiffusion3_5TextToImage", + }, + "ImageToImage": { + "stable_diffusion_xl_refiner": "StableDiffusionXLRefinerImageToImage", }, "TokenClassify": { "bert": "BertTokenClassify", @@ -655,6 +668,12 @@ "speech_to_text_audio": "Speech2TextAudioConfig", "speech_to_text_text": "Speech2TextTextConfig", "stable_diffusion": "StableDiffusionConfig", + "stable_diffusion_2": "StableDiffusion2Config", + "stable_diffusion_xl": "StableDiffusionXLConfig", + "stable_diffusion_xl_refiner": "StableDiffusionXLRefinerConfig", + "stable_diffusion_3": "StableDiffusion3Config", + "stable_diffusion_3_5": "StableDiffusion3_5Config", + "stable_diffusion_3_t5_encoder": "StableDiffusion3T5EncoderConfig", "swin": "SwinConfig", "swinv2": "SwinV2Config", "t5": "T5Config", @@ -750,6 +769,11 @@ "siglip2": "SigLIP2Tokenizer", "speech_to_text": "Speech2TextTokenizer", "stable_diffusion": "StableDiffusionTokenizer", + "stable_diffusion_2": "StableDiffusion2Tokenizer", + "stable_diffusion_xl": "StableDiffusionXLTokenizer", + "stable_diffusion_xl_refiner": "StableDiffusionXLTokenizer", + "stable_diffusion_3": "StableDiffusion3Tokenizer", + "stable_diffusion_3_5": "StableDiffusion3_5Tokenizer", "t5": "T5Tokenizer", "tipsv2": "Tipsv2Tokenizer", "whisper": "WhisperTokenizer", diff --git a/zeromodels/base/base_attention.py b/zeromodels/base/base_attention.py index cc275fd3..ff7b1d18 100644 --- a/zeromodels/base/base_attention.py +++ b/zeromodels/base/base_attention.py @@ -5,30 +5,32 @@ import keras from keras import ops -VALID_ATTN_IMPL = ("sdpa", "flash") +VALID_ATTN_IMPL = ("sdpa", "fused", "flash") DEFAULT_ATTN_IMPLEMENTATION = "sdpa" -# The attention implementation active for the current build or forward. This is -# a ContextVar, not a plain module global: ``Model.from_weights`` activates a -# model's implementation only for the scope of that model's build (so a -# functional graph bakes the right branch) and each of that model's forwards / -# generation steps (so an imperative decode reads it), always restoring the -# previous value. A per-model choice therefore never leaks to another model, a -# later forward, or another thread. -_ATTN_IMPLEMENTATION = contextvars.ContextVar( - "attn_implementation", default=DEFAULT_ATTN_IMPLEMENTATION -) +# The attention implementation explicitly chosen for the current build or +# forward (``None``: nothing chosen, each layer falls back to its own default, +# ``"sdpa"`` unless a model family says otherwise). This is a ContextVar, not a +# plain module global: ``Model.from_weights`` activates a model's implementation +# only for the scope of that model's build (so a functional graph bakes the +# right branch) and each of that model's forwards / generation steps (so an +# imperative decode reads it), always restoring the previous value. A per-model +# choice therefore never leaks to another model, a later forward, or another +# thread. +_ATTN_IMPLEMENTATION = contextvars.ContextVar("attn_implementation", default=None) @contextlib.contextmanager def use_attn_implementation(attn_implementation): """Activate ``attn_implementation`` for the enclosed build or forward. - ``None`` selects the default (``"sdpa"``). The previous value is always - restored on exit, so a per-model implementation never leaks past its scope. + ``None`` clears the choice (layers use their default, ``"sdpa"`` unless the + model family picks another, such as the SD 3 MMDiT's ``"fused"``). The + previous value is always restored on exit, so a per-model implementation + never leaks past its scope. """ - impl = attn_implementation or DEFAULT_ATTN_IMPLEMENTATION - if impl not in VALID_ATTN_IMPL: + impl = attn_implementation + if impl is not None and impl not in VALID_ATTN_IMPL: raise ValueError( f"attn_implementation must be one of {VALID_ATTN_IMPL}, got {impl!r}" ) @@ -39,6 +41,13 @@ def use_attn_implementation(attn_implementation): _ATTN_IMPLEMENTATION.reset(token) +def active_attn_implementation(): + """The implementation explicitly activated for the current build or forward + (by ``Model.from_weights`` or :func:`use_attn_implementation`), or ``None``. + """ + return _ATTN_IMPLEMENTATION.get() + + def with_model_attn_implementation(method): """Decorate a model method so it runs under the model's captured implementation. @@ -97,6 +106,12 @@ def fused_attention( * ``"sdpa"`` -- hand-written matmul/softmax math. Portable across every backend, dtype and device. This is the default. + * ``"fused"`` -- :func:`keras.ops.dot_product_attention` with the backend's + own kernel selection: torch's ``scaled_dot_product_attention`` (flash / + memory-efficient / math, so the logits are never materialized on a GPU + when a fused kernel applies; masks allowed), XLA on JAX. Falls back to + the ``"sdpa"`` math on TensorFlow, with attention dropout or a logit + soft-cap. * ``"flash"`` -- :func:`keras.ops.dot_product_attention` with ``flash_attention=True`` (the real flash kernel). Used only when the backend supports it and there is no attention dropout or logit soft-cap; @@ -120,13 +135,16 @@ def fused_attention( probabilities. Only active during training with a positive rate, in which case the ``"sdpa"`` path is used so it can be applied. training: Whether the call is in training mode. - attn_implementation: ``"sdpa"`` / ``"flash"`` / ``None`` (use the - implementation active in the current context). + attn_implementation: ``"sdpa"`` / ``"fused"`` / ``"flash"`` / ``None`` + (use the implementation active in the current context, else + ``"sdpa"``). Returns: ``(batch, num_heads, q_len, head_dim)``. """ - impl = attn_implementation or _ATTN_IMPLEMENTATION.get() + impl = ( + attn_implementation or _ATTN_IMPLEMENTATION.get() or DEFAULT_ATTN_IMPLEMENTATION + ) if impl not in VALID_ATTN_IMPL: raise ValueError( f"attn_implementation must be one of {VALID_ATTN_IMPL}, got {impl!r}" @@ -135,8 +153,8 @@ def fused_attention( use_dropout = ( bool(training) and dropout is not None and getattr(dropout, "rate", 0.0) > 0.0 ) - use_flash = ( - impl == "flash" + use_fused_op = ( + impl in ("fused", "flash") and _fused_op_available() and not use_dropout and soft_cap is None @@ -145,7 +163,7 @@ def fused_attention( # float32 mask against bf16 logits, but tensorflow raises on the mismatch. if attention_mask is not None: attention_mask = ops.cast(attention_mask, query.dtype) - if use_flash: + if use_fused_op: q = ops.transpose(query, (0, 2, 1, 3)) k = ops.transpose(key, (0, 2, 1, 3)) v = ops.transpose(value, (0, 2, 1, 3)) @@ -155,7 +173,7 @@ def fused_attention( v, bias=attention_mask, scale=scale, - flash_attention=True, + flash_attention=True if impl == "flash" else None, ) return ops.transpose(out, (0, 2, 1, 3)) diff --git a/zeromodels/base/base_diffusion.py b/zeromodels/base/base_diffusion.py index d9ec7221..65515713 100644 --- a/zeromodels/base/base_diffusion.py +++ b/zeromodels/base/base_diffusion.py @@ -28,16 +28,39 @@ class BaseDiffusion: A model plugs in five hooks: - - ``encode_prompt(input_ids, attention_mask=None) -> (batch, seq, dim)`` -- the - conditioning for a batch of token ids (the text tower). + - ``encode_prompt(input_ids, attention_mask=None, **conditioning)`` -- the + conditioning for a batch of token ids (the text tower): a ``(batch, seq, dim)`` + tensor, or any nested structure of batch-major tensors when the denoiser + takes more than one (SDXL's ``{"encoder_hidden_states", "text_embeds", + "time_ids"}``). Extra ``generate`` keyword arguments arrive here (a model's + micro-conditioning), see below. - ``unconditional_ids(batch) -> (batch, seq)`` -- token ids of the empty prompt, - the unconditional branch of classifier-free guidance. + the unconditional branch of classifier-free guidance. A model whose + unconditional branch is not an encoded prompt (SDXL zeroes it) overrides + ``encode_negative_prompt`` instead. - ``predict_noise(latents, timesteps, embeddings) -> (batch, h, w, c)`` -- one - denoiser call (the UNet), ``timesteps`` being ``(batch,)`` floats. + denoiser call (the UNet), ``timesteps`` being ``(batch,)`` floats and + ``embeddings`` whatever ``encode_prompt`` returned (batched ``[uncond, cond]`` + along axis 0 under guidance). - ``decode_latents(latents) -> (batch, H, W, 3)`` -- the decoded image in ``[-1, 1]``, channels-last (the VAE decoder, latent scaling included). - ``latent_shape -> (height, width, channels)`` of one initial latent, and a ``scheduler`` attribute (a :class:`~zeromodels.base.base_scheduler.BaseScheduler`). + - ``encode_latents(image) -> (batch, h, w, c)`` -- optional, the VAE-encoded + (scaled) latent of a ``[-1, 1]`` channels-last image, for image-to-image. + + ``generate`` is also the image-to-image entry point (diffusers' + ``Img2ImgPipeline``): given an ``image`` (or a clean ``latents``) and a + ``strength``, it noises the encoded image to the matching point of the schedule + and denoises the remaining steps. ``denoising_end`` stops a run early and + ``denoising_start`` resumes one from a partially denoised latent, the SDXL + base + refiner "ensemble of experts" (``output_type="latent"`` hands the latent + over). + + Keyword arguments ``generate`` does not know are conditioning and go to + ``encode_prompt``; a ``negative_``-prefixed twin (``negative_original_size`` for + ``original_size``) goes to the negative branch instead, which otherwise reuses + the positive value, the way ``negative_input_ids`` pairs with ``input_ids``. Generation settings resolve like the LM mixins': an explicit ``generate`` argument wins, then the instance's ``generate_args`` (a repo's ``zm_config.json`` @@ -49,6 +72,7 @@ class BaseDiffusion: DEFAULT_NUM_INFERENCE_STEPS = 50 DEFAULT_GUIDANCE_SCALE = 7.5 + DEFAULT_STRENGTH = 0.8 def encode_prompt(self, input_ids, attention_mask=None): raise NotImplementedError( @@ -60,6 +84,20 @@ def unconditional_ids(self, batch): f"{type(self).__name__} must implement unconditional_ids()." ) + def encode_negative_prompt(self, negative_input_ids, batch, **conditioning): + """The unconditional branch of classifier-free guidance: the encoded negative + prompt, or the encoded empty prompt (``unconditional_ids``) when none is + given. A single ``(1, seq)`` negative row is shared by the whole batch. + """ + if negative_input_ids is None: + negative_input_ids = self.unconditional_ids(batch) + negative_input_ids = ops.cast( + ops.convert_to_tensor(negative_input_ids), "int32" + ) + if int(negative_input_ids.shape[0]) == 1 and batch > 1: + negative_input_ids = ops.repeat(negative_input_ids, batch, axis=0) + return self.encode_prompt(negative_input_ids, **conditioning) + def predict_noise(self, latents, timesteps, embeddings): raise NotImplementedError( f"{type(self).__name__} must implement predict_noise()." @@ -74,6 +112,11 @@ def decode_latents(self, latents): def latent_shape(self): raise NotImplementedError(f"{type(self).__name__} must define latent_shape.") + def encode_latents(self, image): + raise NotImplementedError( + f"{type(self).__name__} must implement encode_latents() for image-to-image." + ) + def generate( self, input_ids, @@ -83,8 +126,15 @@ def generate( guidance_scale=None, seed=None, latents=None, + image=None, + strength=None, + denoising_start=None, + denoising_end=None, + output_type="image", + **conditioning, ): - """Generate images from tokenized prompts. + """Generate images from tokenized prompts (text-to-image), or from an image + or latent and a prompt (image-to-image). Args: input_ids: ``(batch, seq)`` token ids, i.e. ``**tokenizer(prompts)``. @@ -95,38 +145,136 @@ def generate( num_inference_steps: Scheduler steps (``generate_args`` / 50). guidance_scale: Classifier-free guidance strength (``generate_args`` / 7.5); ``<= 1`` disables it. - seed: RNG seed for the initial latent (reproducible per backend). - latents: Explicit initial latent ``(batch, *latent_shape)``, for results - that are identical across backends. + seed: RNG seed for the initial latent / added noise (reproducible per + backend). + latents: Text-to-image: the explicit initial noise + ``(batch, *latent_shape)``, for results identical across backends. + Image-to-image (with ``strength`` or ``denoising_start``): the + clean, scaled latent to start from instead of ``image``. + image: ``(batch, H, W, 3)`` uint8 or ``[0, 1]`` float images to start + from (image-to-image); encoded by the VAE, then noised to the + ``strength`` point of the schedule. + strength: Image-to-image: the fraction of the schedule to run, + ``(0, 1]`` (``1.0`` ignores the image); the noise added matches. + Defaults to ``generate_args["strength"]`` when an ``image`` or + ``latents`` are given without ``denoising_start``, else 0.8. + denoising_start: Resume from a partially denoised latent at this + fraction of the schedule, adding no noise (the refiner's half of + the SDXL ensemble: the base ran with ``denoising_end`` at the + same value and ``output_type="latent"``). + denoising_end: Stop after this fraction of the schedule (the base's + half of the ensemble). + output_type: ``"image"`` (uint8 images) or ``"latent"`` (the final + latent, a backend tensor, e.g. for a refiner). + **conditioning: Model-specific conditioning handed to ``encode_prompt`` + (SDXL's ``original_size`` / ``crops_coords_top_left`` / + ``target_size``); ``negative_`` variants apply to the negative + branch only. Returns: - ``(batch, H, W, 3)`` uint8 numpy images. + ``(batch, H, W, 3)`` uint8 numpy images, or the latent. """ num_inference_steps, guidance_scale, seed = self.resolve_generation_args( num_inference_steps, guidance_scale, seed ) + positive, negative = self.split_conditioning(conditioning) input_ids = ops.cast(ops.convert_to_tensor(input_ids), "int32") batch = int(input_ids.shape[0]) + image_to_image = ( + image is not None or strength is not None or denoising_start is not None + ) + if image_to_image and strength is None and denoising_start is None: + strength = (getattr(self, "generate_args", None) or {}).get( + "strength", self.DEFAULT_STRENGTH + ) with inference_scope(): - embeddings = self.encode_prompt(input_ids, attention_mask) + embeddings = self.encode_prompt(input_ids, attention_mask, **positive) if guidance_scale > 1.0: - uncond_ids = ( - ops.cast(ops.convert_to_tensor(negative_input_ids), "int32") - if negative_input_ids is not None - else self.unconditional_ids(batch) + uncond = self.encode_negative_prompt( + negative_input_ids, batch, **negative + ) + embeddings = keras.tree.map_structure( + lambda u, c: ops.concatenate([u, c], axis=0), uncond, embeddings ) - if int(uncond_ids.shape[0]) == 1 and batch > 1: - uncond_ids = ops.repeat(uncond_ids, batch, axis=0) # one for all - embeddings = ops.concatenate( - [self.encode_prompt(uncond_ids), embeddings], axis=0 + # the scheduler's init_noise_sigma depends on the inference timesteps + self.scheduler.set_timesteps(num_inference_steps) + timesteps = self.scheduler.timesteps + if denoising_end is not None: + timesteps = timesteps[: self.steps_before(denoising_end, timesteps)] + if image_to_image: + timesteps, t_start = self.image_to_image_timesteps( + timesteps, strength, denoising_start ) - latents = self.prepare_latents(batch, seed=seed, latents=latents) + latents = self.prepare_image_latents( + image, + latents, + timesteps[0] if denoising_start is None else None, + seed=seed, + ) + self.scheduler.set_begin_index(t_start) + else: + latents = self.prepare_latents(batch, seed=seed, latents=latents) latents = self.denoise( - latents, embeddings, num_inference_steps, guidance_scale + latents, + embeddings, + num_inference_steps, + guidance_scale, + timesteps=timesteps, ) + if output_type == "latent": + return latents image = self.decode_latents(latents) return self.postprocess_image(image) + @staticmethod + def steps_before(fraction, timesteps): + """How many of ``timesteps`` lie in the first ``fraction`` of the schedule + (timesteps at or above the ``1 - fraction`` cutoff, diffusers' rounding).""" + cutoff = int(round(1000 - fraction * 1000)) + return int(np.sum(np.asarray(timesteps, dtype=np.float64) >= cutoff)) + + def image_to_image_timesteps(self, timesteps, strength, denoising_start): + """The tail of ``timesteps`` an image-to-image run covers and its start + index: the last ``strength`` fraction, or everything after + ``denoising_start``.""" + n = len(timesteps) + if denoising_start is not None: + t_start = self.steps_before(denoising_start, timesteps) + else: + t_start = max(n - min(int(n * strength), n), 0) + return timesteps[t_start:], t_start + + def prepare_image_latents(self, image, latents, timestep, seed=None): + """The starting latent of an image-to-image run: the encoded ``image`` (or + the given clean ``latents``), noised to ``timestep`` unless ``timestep`` is + ``None`` (resuming a partially denoised latent).""" + if latents is None: + if image is None: + raise ValueError("Image-to-image needs an image or latents.") + image = ops.convert_to_tensor(image) + if not keras.backend.is_float_dtype(image.dtype): + image = ops.cast(image, "float32") / 255.0 + latents = self.encode_latents(ops.cast(image, "float32") * 2.0 - 1.0) + latents = ops.cast(ops.convert_to_tensor(latents), "float32") + if timestep is None: + return latents + noise = keras.random.normal(ops.shape(latents), seed=seed, dtype="float32") + return self.scheduler.add_noise(latents, noise, [timestep]) + + @staticmethod + def split_conditioning(conditioning): + """Split ``generate``'s extra keyword arguments into the positive branch's + and the negative branch's (``negative_`` overrides ```` there; + a ``None`` override keeps the positive value).""" + positive = { + k: v for k, v in conditioning.items() if not k.startswith("negative_") + } + negative = dict(positive) + for key, value in conditioning.items(): + if key.startswith("negative_") and value is not None: + negative[key[len("negative_") :]] = value + return positive, negative + def resolve_generation_args(self, num_inference_steps, guidance_scale, seed): defaults = getattr(self, "generate_args", None) or {} if num_inference_steps is None: @@ -163,12 +311,16 @@ def classifier_free_guidance(noise_pred, guidance_scale): uncond, cond = ops.split(noise_pred, 2, axis=0) return uncond + guidance_scale * (cond - uncond) - def denoise(self, latents, embeddings, num_inference_steps, guidance_scale): + def denoise( + self, latents, embeddings, num_inference_steps, guidance_scale, timesteps=None + ): """Run the scheduler's denoising loop over ``predict_noise``. - ``embeddings`` is ``[uncond, cond]`` batched along axis 0 when - ``guidance_scale > 1`` (classifier-free guidance), else just ``cond``. - Returns the denoised latent. + ``embeddings`` is whatever ``encode_prompt`` returned, ``[uncond, cond]`` + batched along axis 0 when ``guidance_scale > 1`` (classifier-free + guidance), else just ``cond``. Loops over the scheduler's + ``num_inference_steps`` timesteps, or over the given ``timesteps`` (a + subset of them, the scheduler already set up). Returns the denoised latent. """ do_cfg = guidance_scale > 1.0 batch = int(ops.shape(latents)[0]) @@ -176,8 +328,10 @@ def denoise(self, latents, embeddings, num_inference_steps, guidance_scale): step = self.cached_denoise_step(do_cfg) guidance = ops.convert_to_tensor(guidance_scale, dtype="float32") scheduler = self.scheduler - scheduler.set_timesteps(num_inference_steps) - for t in scheduler.timesteps: + if timesteps is None: + scheduler.set_timesteps(num_inference_steps) + timesteps = scheduler.timesteps + for t in timesteps: model_input = ( ops.concatenate([latents, latents], axis=0) if do_cfg else latents ) diff --git a/zeromodels/base/base_mixin.py b/zeromodels/base/base_mixin.py index a83e0a04..247dfa14 100644 --- a/zeromodels/base/base_mixin.py +++ b/zeromodels/base/base_mixin.py @@ -264,10 +264,13 @@ def from_weights( backbone. Applied on both the repo-id ``zm_config.json`` load path and the ``hf:`` / variant converter transfer path (mismatched targets left at init). - attn_implementation: ``"sdpa"`` (portable manual math, the default) - or ``"flash"`` (``keras.ops.dot_product_attention`` with the - flash kernel; needs a flash-capable GPU/TPU and fp16/bf16). Set - before the model is built. + attn_implementation: ``"sdpa"`` (portable manual math, the default + of most layers), ``"fused"`` (``keras.ops.dot_product_attention`` + with the backend's own kernel selection: torch's flash / + memory-efficient kernels, XLA on JAX; the SD 3 MMDiT's default) + or ``"flash"`` (the flash kernel; needs a flash-capable GPU/TPU + and fp16/bf16). ``None`` leaves every layer at its own default. + Set before the model is built. quantization: ``None`` (default), ``"int8"``, ``"int4"`` or ``"fp8"`` (or a :class:`~zeromodels.quantization.\ QuantizationConfig` / scheme). When set, the model is quantized weight-only: @@ -314,9 +317,7 @@ def from_weights( f"attn_implementation must be one of " f"{base_attention.VALID_ATTN_IMPL}, got {attn_implementation!r}" ) - resolved_attn = ( - attn_implementation or base_attention.DEFAULT_ATTN_IMPLEMENTATION - ) + resolved_attn = attn_implementation if load_dtype is None: load_dtype = cls.hub_repo_weight_dtype(identifier) diff --git a/zeromodels/base/base_scheduler.py b/zeromodels/base/base_scheduler.py index 1b70e6da..9bebb6c4 100644 --- a/zeromodels/base/base_scheduler.py +++ b/zeromodels/base/base_scheduler.py @@ -126,6 +126,10 @@ def scale_model_input(self, sample, timestep=None): def set_timesteps(self, num_inference_steps): raise NotImplementedError + def set_begin_index(self, index): + """Start the loop at ``self.timesteps[index]`` (image-to-image skips the first + steps); a no-op for the samplers that index by timestep value.""" + def step(self, model_output, timestep, sample, **kwargs): raise NotImplementedError @@ -334,55 +338,126 @@ def to_config(self): class EulerDiscreteScheduler(BaseScheduler): - """Euler sampler over the karras-style sigma parameterization.""" + """Euler sampler over the karras-style sigma parameterization. - def __init__(self, **kwargs): + Args: + timestep_spacing: How the inference timesteps are spread over the training + ones: ``"linspace"`` (evenly from ``num_train_timesteps - 1`` to 0), + ``"leading"`` (multiples of the step ratio from 0, shifted by + ``steps_offset``; the SD 1.x / 2.x and SDXL repos) or ``"trailing"`` + (multiples counted back from ``num_train_timesteps``; SDXL-Turbo). + interpolation_type: ``"linear"`` interpolates the training sigmas at the + timesteps; ``"log_linear"`` spaces them evenly in log space. + """ + + def __init__( + self, timestep_spacing="linspace", interpolation_type="linear", **kwargs + ): super().__init__(**kwargs) + self.timestep_spacing = timestep_spacing + self.interpolation_type = interpolation_type ac = self.alphas_cumprod sigmas = ((np.float32(1.0) - ac) / ac) ** np.float32(0.5) self.train_sigmas = sigmas.astype(np.float32) self.sigmas = np.concatenate([sigmas[::-1], [0.0]]).astype(np.float32) + self.step_index = 0 @property def init_noise_sigma(self): - # "linspace" timestep spacing (the only one implemented): the reference - # scales the initial noise by the largest sigma itself. - return float(self.sigmas.max()) + # the reference scales the initial noise by the largest sigma itself for + # "linspace" / "trailing" spacing and by sqrt(sigma_max^2 + 1) for "leading" + max_sigma = float(self.sigmas.max()) + if self.timestep_spacing in ("linspace", "trailing"): + return max_sigma + return (max_sigma**2 + 1) ** 0.5 def scale_model_input(self, sample, timestep=None): sigma = self.sigmas[self.step_index] return sample / ((sigma**2 + 1) ** 0.5) + def set_begin_index(self, index): + self.step_index = int(index) + + def add_noise(self, original_samples, noise, timesteps): + # the Euler sample space is unscaled: x_t = x_0 + sigma_t * noise, sigma_t + # being the inference sigma of the timestep (set_timesteps first) + timesteps = np.asarray(timesteps, dtype=np.float32).reshape(-1) + index = [int(np.nonzero(self.timesteps == t)[0][0]) for t in timesteps] + sigma = self.sigmas[index] + while sigma.ndim < len(ops.shape(original_samples)): + sigma = sigma[..., None] + sigma = ops.convert_to_tensor(sigma, dtype=original_samples.dtype) + return original_samples + noise * sigma + + def spaced_timesteps(self, num_inference_steps): + n_train = self.num_train_timesteps + if self.timestep_spacing == "linspace": + timesteps = np.linspace( + 0, n_train - 1, num_inference_steps, dtype=np.float32 + ) + return timesteps[::-1].copy() + if self.timestep_spacing == "leading": + step_ratio = n_train // num_inference_steps + timesteps = (np.arange(0, num_inference_steps) * step_ratio).round() + return timesteps[::-1].copy().astype(np.float32) + self.steps_offset + if self.timestep_spacing == "trailing": + step_ratio = n_train / num_inference_steps + timesteps = np.arange(n_train, 0, -step_ratio).round() + return timesteps.astype(np.float32) - 1 + raise ValueError( + f"Unknown timestep_spacing {self.timestep_spacing!r}; expected " + "'linspace', 'leading' or 'trailing'." + ) + def set_timesteps(self, num_inference_steps): self.num_inference_steps = num_inference_steps - timesteps = np.linspace( - 0, self.num_train_timesteps - 1, num_inference_steps, dtype=np.float32 - )[::-1].copy() + timesteps = self.spaced_timesteps(num_inference_steps) sigmas = self.train_sigmas - interp = np.interp(timesteps, np.arange(len(sigmas)), sigmas) - self.sigmas = np.concatenate([interp, [0.0]]).astype(np.float32) + if self.interpolation_type == "linear": + sigmas = np.interp(timesteps, np.arange(len(sigmas)), sigmas) + elif self.interpolation_type == "log_linear": + sigmas = np.exp( + np.linspace( + np.log(sigmas[-1]), np.log(sigmas[0]), num_inference_steps + 1 + ) + )[:-1] + else: + raise ValueError( + f"Unknown interpolation_type {self.interpolation_type!r}; expected " + "'linear' or 'log_linear'." + ) + self.sigmas = np.concatenate([sigmas, [0.0]]).astype(np.float32) self.timesteps = timesteps self.step_index = 0 return self.timesteps - def step(self, model_output, timestep, sample, **kwargs): - sigma = float(self.sigmas[self.step_index]) - sigma_next = float(self.sigmas[self.step_index + 1]) + def pred_original_sample(self, model_output, sigma, sample): + # float32 sigma arithmetic, so the coefficients match the reference's bits if self.prediction_type == "epsilon": - pred_original = sample - sigma * model_output - elif self.prediction_type == "v_prediction": - pred_original = model_output * (-sigma / (sigma**2 + 1) ** 0.5) + ( + return sample - sigma * model_output + if self.prediction_type == "v_prediction": + return model_output * (-sigma / (sigma**2 + 1) ** 0.5) + ( sample / (sigma**2 + 1) ) - elif self.prediction_type == "sample": - pred_original = model_output - else: - raise ValueError(f"Unknown prediction_type {self.prediction_type!r}.") + if self.prediction_type == "sample": + return model_output + raise ValueError(f"Unknown prediction_type {self.prediction_type!r}.") + + def step(self, model_output, timestep, sample, **kwargs): + sigma = self.sigmas[self.step_index] + sigma_next = self.sigmas[self.step_index + 1] + pred_original = self.pred_original_sample(model_output, sigma, sample) derivative = (sample - pred_original) / sigma prev_sample = sample + derivative * (sigma_next - sigma) self.step_index += 1 return prev_sample + def to_config(self): + config = super().to_config() + config["timestep_spacing"] = self.timestep_spacing + config["interpolation_type"] = self.interpolation_type + return config + class EulerAncestralDiscreteScheduler(EulerDiscreteScheduler): """Ancestral Euler sampler: injects fresh noise at each step (stochastic).""" @@ -392,23 +467,15 @@ def __init__(self, seed=None, **kwargs): self.seed_generator = keras_random.SeedGenerator(seed) def step(self, model_output, timestep, sample, **kwargs): - sigma = float(self.sigmas[self.step_index]) - sigma_next = float(self.sigmas[self.step_index + 1]) - if self.prediction_type == "epsilon": - pred_original = sample - sigma * model_output - elif self.prediction_type == "v_prediction": - pred_original = model_output * (-sigma / (sigma**2 + 1) ** 0.5) + ( - sample / (sigma**2 + 1) - ) - else: - raise ValueError(f"Unknown prediction_type {self.prediction_type!r}.") - sigma_up = min( - sigma_next, + sigma = self.sigmas[self.step_index] + sigma_next = self.sigmas[self.step_index + 1] + pred_original = self.pred_original_sample(model_output, sigma, sample) + sigma_up = ( (sigma_next**2 * (sigma**2 - sigma_next**2) / sigma**2) ** 0.5 if sigma > 0 - else 0.0, + else np.float32(0.0) ) - sigma_down = (max(sigma_next**2 - sigma_up**2, 0.0)) ** 0.5 + sigma_down = (sigma_next**2 - sigma_up**2) ** 0.5 derivative = (sample - pred_original) / sigma prev_sample = sample + derivative * (sigma_down - sigma) noise = keras_random.normal( @@ -419,11 +486,92 @@ def step(self, model_output, timestep, sample, **kwargs): return prev_sample +class FlowMatchEulerDiscreteScheduler(BaseScheduler): + """Euler sampler for rectified-flow models (Stable Diffusion 3 / 3.5, FLUX). + + There is no beta schedule: the noise level is the flow time ``sigma`` in + ``[0, 1]`` (``x_t = (1 - sigma) x_0 + sigma * noise``), the model predicts the + velocity and a step is ``x + (sigma_next - sigma) * v``. The timesteps handed to + the model are ``sigma * num_train_timesteps``. + + Args: + num_train_timesteps: The flow time resolution (1000). + shift: Timestep shift towards noisier levels, ``shift * s / (1 + (shift - 1) s)`` + (3.0 for SD3 / SD3.5). + """ + + def __init__(self, num_train_timesteps=1000, shift=1.0, **kwargs): + # the base (beta) schedule is irrelevant here; keep the constructor + # compatible with from_config (a repo's scheduler_config may carry extra keys) + super().__init__(num_train_timesteps=num_train_timesteps) + self.shift = shift + timesteps = np.linspace( + 1, num_train_timesteps, num_train_timesteps, dtype=np.float32 + )[::-1].copy() + sigmas = timesteps / np.float32(num_train_timesteps) + sigmas = self.shift_sigmas(sigmas) + self.sigma_min = float(sigmas[-1]) + self.sigma_max = float(sigmas[0]) + self.sigmas = sigmas + self.timesteps = sigmas * np.float32(num_train_timesteps) + self.step_index = 0 + + def shift_sigmas(self, sigmas): + shift = np.float32(self.shift) + return (shift * sigmas / (1 + (shift - 1) * sigmas)).astype(np.float32) + + @property + def init_noise_sigma(self): + return 1.0 + + def set_begin_index(self, index): + self.step_index = int(index) + + def set_timesteps(self, num_inference_steps): + self.num_inference_steps = num_inference_steps + n_train = self.num_train_timesteps + timesteps = np.linspace( + self.sigma_max * n_train, self.sigma_min * n_train, num_inference_steps + ) + sigmas = timesteps / n_train + sigmas = self.shift * sigmas / (1 + (self.shift - 1) * sigmas) + sigmas = sigmas.astype(np.float32) + self.sigmas = np.concatenate([sigmas, [0.0]]).astype(np.float32) + self.timesteps = sigmas * np.float32(n_train) + self.step_index = 0 + return self.timesteps + + def add_noise(self, original_samples, noise, timesteps): + # the forward flow: x_t = (1 - sigma) x_0 + sigma * noise + timesteps = np.asarray(timesteps, dtype=np.float32).reshape(-1) + index = [int(np.nonzero(self.timesteps == t)[0][0]) for t in timesteps] + sigma = self.sigmas[index] + while sigma.ndim < len(ops.shape(original_samples)): + sigma = sigma[..., None] + sigma = ops.convert_to_tensor(sigma, dtype=original_samples.dtype) + return sigma * noise + (1.0 - sigma) * original_samples + + def step(self, model_output, timestep, sample, **kwargs): + sigma = self.sigmas[self.step_index] + sigma_next = self.sigmas[self.step_index + 1] + prev_sample = sample + (sigma_next - sigma) * model_output + self.step_index += 1 + return prev_sample + + def to_config(self): + return { + "_class_name": type(self).__name__, + "num_train_timesteps": self.num_train_timesteps, + "shift": self.shift, + } + + SCHEDULER_REGISTRY = { "DDIMScheduler": DDIMScheduler, "PNDMScheduler": PNDMScheduler, "EulerDiscreteScheduler": EulerDiscreteScheduler, "EulerAncestralDiscreteScheduler": EulerAncestralDiscreteScheduler, + "FlowMatchEulerDiscreteScheduler": FlowMatchEulerDiscreteScheduler, } diff --git a/zeromodels/models/__init__.py b/zeromodels/models/__init__.py index a08ae0c7..4707d7a9 100644 --- a/zeromodels/models/__init__.py +++ b/zeromodels/models/__init__.py @@ -121,6 +121,10 @@ siglip2, speech2text, stable_diffusion, + stable_diffusion_2, + stable_diffusion_3, + stable_diffusion_3_5, + stable_diffusion_xl, swin, swinv2, t5, diff --git a/zeromodels/models/clip/clip_model.py b/zeromodels/models/clip/clip_model.py index 5a77594c..87324fd2 100644 --- a/zeromodels/models/clip/clip_model.py +++ b/zeromodels/models/clip/clip_model.py @@ -240,7 +240,7 @@ def clip_text_backbone( )(encoded_output) indices = ops.argmax(inputs, axis=-1) - one_hot_indices = ops.one_hot(indices, max_seq_len) + one_hot_indices = ops.one_hot(indices, max_seq_len, dtype=last_hidden_state.dtype) pooler_output = ops.einsum("bi,bij->bj", one_hot_indices, last_hidden_state) return last_hidden_state, pooler_output @@ -508,15 +508,21 @@ def __init__( if isinstance(input_tensor, dict): token_ids_input = input_tensor.get("token_ids") if token_ids_input is None: - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) padding_mask_input = input_tensor.get("padding_mask") if padding_mask_input is None: padding_mask_input = layers.Input( - shape=[max_seq_len], name="padding_mask" + shape=[max_seq_len], dtype="int32", name="padding_mask" ) else: - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") - padding_mask_input = layers.Input(shape=[max_seq_len], name="padding_mask") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) + padding_mask_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="padding_mask" + ) last_hidden_state, pooler_output = clip_text_backbone( token_ids_input, @@ -837,15 +843,21 @@ def __init__( if isinstance(input_tensor, dict): token_ids_input = input_tensor.get("token_ids") if token_ids_input is None: - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) padding_mask_input = input_tensor.get("padding_mask") if padding_mask_input is None: padding_mask_input = layers.Input( - shape=[max_seq_len], name="padding_mask" + shape=[max_seq_len], dtype="int32", name="padding_mask" ) else: - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") - padding_mask_input = layers.Input(shape=[max_seq_len], name="padding_mask") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) + padding_mask_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="padding_mask" + ) text_model = CLIPTextModel( max_seq_len=max_seq_len, @@ -1034,16 +1046,22 @@ def __init__( images_input = layers.Input(shape=input_shape, name="images") token_ids_input = input_tensor.get("token_ids") if token_ids_input is None: - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) padding_mask_input = input_tensor.get("padding_mask") if padding_mask_input is None: padding_mask_input = layers.Input( - shape=[max_seq_len], name="padding_mask" + shape=[max_seq_len], dtype="int32", name="padding_mask" ) else: images_input = layers.Input(shape=input_shape, name="images") - token_ids_input = layers.Input(shape=[max_seq_len], name="token_ids") - padding_mask_input = layers.Input(shape=[max_seq_len], name="padding_mask") + token_ids_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="token_ids" + ) + padding_mask_input = layers.Input( + shape=[max_seq_len], dtype="int32", name="padding_mask" + ) vision_model = CLIPVisionModel( image_size=image_size, diff --git a/zeromodels/models/stable_diffusion/convert_stable_diffusion_diffusers_to_keras.py b/zeromodels/models/stable_diffusion/convert_stable_diffusion_diffusers_to_keras.py index 816d73c6..84a68e64 100644 --- a/zeromodels/models/stable_diffusion/convert_stable_diffusion_diffusers_to_keras.py +++ b/zeromodels/models/stable_diffusion/convert_stable_diffusion_diffusers_to_keras.py @@ -1,4 +1,4 @@ -import collections.abc +from typing import Dict import numpy as np from tqdm import tqdm @@ -10,7 +10,6 @@ from zeromodels.conversion.weight_split_util import split_model_weights from zeromodels.conversion.weight_transfer_util import ( compare_keras_torch_names, - transfer_attention_weights, transfer_weights, ) @@ -22,154 +21,48 @@ "stable-diffusion-v1-5": "stable-diffusion-v1-5/stable-diffusion-v1-5", } -LEGACY_ATTENTION_KEYS = { - ".query.": ".to_q.", - ".key.": ".to_k.", - ".value.": ".to_v.", - ".proj_attn.": ".to_out.0.", +WEIGHT_NAME_MAPPING: Dict[str, str] = { + "__": ".", + "/kernel": ".weight", + "/gamma": ".weight", + "/beta": ".bias", + "/scale": ".weight", # RMSNormalization (the SD 3.5 qk norms) + "/": ".", } -class RenamedStateDict(collections.abc.Mapping): - def __init__(self, state_dict): - self.state_dict = state_dict - self.to_source = {} - for key in state_dict: - new = key - if ".attentions." in key: - for old, current in LEGACY_ATTENTION_KEYS.items(): - new = new.replace(old, current) - self.to_source[new] = key - - def __getitem__(self, key): - return self.state_dict[self.to_source[key]] - - def __contains__(self, key): - return key in self.to_source - - def __iter__(self): - return iter(self.to_source) - - def __len__(self): - return len(self.to_source) - - -WEIGHT_SUFFIX = {"kernel": "weight", "bias": "bias", "gamma": "weight", "beta": "bias"} - -ATTN_NAME_REPLACE = { - "..": ".", - "down.blocks": "down_blocks", - "up.blocks": "up_blocks", - "mid.block": "mid_block", - "transformer.blocks": "transformer_blocks", - "proj.in": "proj_in", - "proj.out": "proj_out", - "to.q": "to_q", - "to.k": "to_k", - "to.v": "to_v", - "to.out": "to_out", - "group.norm": "group_norm", -} - - -def torch_key(keras_weight): - leaf, variable = keras_weight.path.split("/")[-2:] - return f"{leaf.replace('__', '.')}.{WEIGHT_SUFFIX[variable]}" - - -def transfer_component(keras_model, state): - consumed = set() - trainable, non_trainable = split_model_weights(keras_model) - - for keras_weight, _ in tqdm( - trainable + non_trainable, desc="Transferring weights to Keras" - ): - key = torch_key(keras_weight) - consumed.add(key) - - if "attention" in key: - transfer_attention_weights( - keras_weight.path, keras_weight, state, ATTN_NAME_REPLACE - ) - continue - - if key not in state: - raise WeightMappingError(keras_weight.path, key) - torch_weight = state[key] - if not compare_keras_torch_names( - keras_weight.path, keras_weight, key, torch_weight - ): - raise WeightShapeMismatchError( - keras_weight.path, keras_weight.shape, key, torch_weight.shape - ) - if key.startswith("time_embedding.") and keras_weight.ndim == 2: - # transfer_weights treats any "embedding" 2D weight as a lookup table; - # the timestep MLP's linears are plain dense kernels - keras_weight.assign(np.transpose(torch_weight)) - continue - transfer_weights(key, keras_weight, torch_weight) - - unused = sorted(set(state) - consumed) - if unused: - raise ValueError( - f"{type(keras_model).__name__}: {len(unused)} checkpoint tensors unused, " - f"e.g. {unused[:3]}." - ) - - -def numpy_state_dict(module): - return {k: v.detach().cpu().numpy() for k, v in module.state_dict().items()} - - -def diffusers_configs(repo, token=None): +def config_from_diffusers(repo, token=None, config_cls=None): from diffusers import AutoencoderKL, PNDMScheduler, UNet2DConditionModel - from transformers import CLIPTextConfig - - return { - "unet": dict( - UNet2DConditionModel.load_config(repo, subfolder="unet", token=token) - ), - "vae": dict(AutoencoderKL.load_config(repo, subfolder="vae", token=token)), - "text": CLIPTextConfig.from_pretrained( - repo, subfolder="text_encoder", token=token - ).to_dict(), - # load_config only reads the json, whichever scheduler the repo declares - "scheduler": dict( - PNDMScheduler.load_config(repo, subfolder="scheduler", token=token) - ), - } - + from transformers import CLIPTextConfig, CLIPTokenizerFast -def config_from_diffusers(repo, token=None): from zeromodels.models.stable_diffusion.stable_diffusion_config import ( StableDiffusionConfig, ) + from zeromodels.models.stable_diffusion.stable_diffusion_model import ( + UNet2DConditionModel as KerasUNet2DConditionModel, + ) - src = diffusers_configs(repo, token=token) - unet, vae, text = src["unet"], src["vae"], src["text"] + config_cls = config_cls or StableDiffusionConfig + unet = dict(UNet2DConditionModel.load_config(repo, subfolder="unet", token=token)) + vae = dict(AutoencoderKL.load_config(repo, subfolder="vae", token=token)) + text = CLIPTextConfig.from_pretrained( + repo, subfolder="text_encoder", token=token + ).to_dict() + # load_config only reads the json, whichever scheduler the repo declares scheduler = { k: v - for k, v in src["scheduler"].items() + for k, v in PNDMScheduler.load_config( + repo, subfolder="scheduler", token=token + ).items() if k == "_class_name" or not k.startswith("_") } - - heads = unet.get("num_attention_heads") or unet.get("attention_head_dim", 8) - if isinstance(heads, (list, tuple)): - heads = heads[0] + # the pad token differs between checkpoints (<|endoftext|> 49407, "!" 0) + tokenizer = CLIPTokenizerFast.from_pretrained( + repo, subfolder="tokenizer", token=token + ) hidden = text["hidden_size"] - return StableDiffusionConfig( - unet_config={ - "sample_size": unet.get("sample_size", 64), - "in_channels": unet.get("in_channels", 4), - "out_channels": unet.get("out_channels", 4), - "down_block_types": tuple(unet["down_block_types"]), - "up_block_types": tuple(unet["up_block_types"]), - "block_out_channels": tuple(unet["block_out_channels"]), - "layers_per_block": unet.get("layers_per_block", 2), - "cross_attention_dim": unet.get("cross_attention_dim", 768), - "num_attention_heads": heads, - "norm_num_groups": unet.get("norm_num_groups", 32), - }, + return config_cls( + unet_config=KerasUNet2DConditionModel.kwargs_from_diffusers_config(unet), vae_config={ "in_channels": vae.get("in_channels", 3), "out_channels": vae.get("out_channels", 3), @@ -177,8 +70,14 @@ def config_from_diffusers(repo, token=None): "block_out_channels": tuple(vae["block_out_channels"]), "layers_per_block": vae.get("layers_per_block", 2), "norm_num_groups": vae.get("norm_num_groups", 32), - "sample_size": vae.get("sample_size", 512), - "scaling_factor": vae.get("scaling_factor", 0.18215), + # the pipeline resolution is the UNet's latent size x the VAE's + # compression; the VAE config's own sample_size is not what the + # checkpoint generates at (768 on the 512px SD 2.1-base repo) + "sample_size": unet.get("sample_size", 64) + * 2 ** (len(vae["block_out_channels"]) - 1), + "scaling_factor": vae.get("scaling_factor") + or 0.18215, # null in some repos + "force_upcast": bool(vae.get("force_upcast", False)), }, text_config={ "hidden_dim": hidden, @@ -191,10 +90,13 @@ def config_from_diffusers(repo, token=None): hidden_act=text.get("hidden_act", "quick_gelu"), layer_norm_eps=text.get("layer_norm_eps", 1e-5), scheduler_config=scheduler, + bos_token_id=tokenizer.bos_token_id, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, ) -def build_from_diffusers(repo, token=None): +def transfer_stable_diffusion(repo, token=None, model_cls=None, config_cls=None): import gc import torch @@ -205,29 +107,75 @@ def build_from_diffusers(repo, token=None): from zeromodels.models.clip import CLIPTextModel from zeromodels.models.stable_diffusion.stable_diffusion_model import ( StableDiffusionModel, - UNet2DConditionModel, ) - config = config_from_diffusers(repo, token=token) - model = StableDiffusionModel(config) + model_cls = model_cls or StableDiffusionModel + config = config_from_diffusers(repo, token=token, config_cls=config_cls) + model = model_cls(config) load = {"torch_dtype": torch.float32, "token": token} - unet = DiffusersUNet2DConditionModel.from_pretrained(repo, subfolder="unet", **load) - UNet2DConditionModel.transfer_from_hf(model.unet, numpy_state_dict(unet)) - del unet - gc.collect() - - vae = DiffusersAutoencoderKL.from_pretrained(repo, subfolder="vae", **load) - model.vae.transfer_from_hf(numpy_state_dict(vae)) - del vae - gc.collect() + # the SD 1.x era VAE files spell the mid-block attention query / key / value / + # proj_attn (diffusers renames them on load, a raw checkpoint read keeps them): + # map them to the to_q / to_k / to_v / to_out.0 names the layers use + legacy = { + ".query.": ".to_q.", + ".key.": ".to_k.", + ".value.": ".to_v.", + ".proj_attn.": ".to_out.0.", + } + for component, module_cls, subfolder in ( + (model.unet, DiffusersUNet2DConditionModel, "unet"), + (model.vae, DiffusersAutoencoderKL, "vae"), + ): + module = module_cls.from_pretrained(repo, subfolder=subfolder, **load) + state = {} + for key, value in module.state_dict().items(): + if subfolder == "vae" and ".attentions." in key: + for old, new in legacy.items(): + key = key.replace(old, new) + state[key] = value.detach().cpu().numpy() + del module + consumed = set() + trainable, non_trainable = split_model_weights(component) + for keras_weight, _ in tqdm( + trainable + non_trainable, desc=f"Transferring {subfolder} weights to Keras" + ): + key = "/".join(keras_weight.path.split("/")[-2:]) + for old, new in WEIGHT_NAME_MAPPING.items(): + key = key.replace(old, new) + consumed.add(key) + if key not in state: + raise WeightMappingError(keras_weight.path, key) + torch_weight = state[key] + if not compare_keras_torch_names( + keras_weight.path, keras_weight, key, torch_weight + ): + raise WeightShapeMismatchError( + keras_weight.path, keras_weight.shape, key, torch_weight.shape + ) + if key.startswith("time_embedding.") and keras_weight.ndim == 2: + # transfer_weights treats any "embedding" 2D weight as a lookup + # table; the timestep MLP's linears are dense kernels + keras_weight.assign(np.transpose(torch_weight)) + continue + transfer_weights(key, keras_weight, torch_weight) + unused = sorted(set(state) - consumed) + if unused: + raise ValueError( + f"{type(component).__name__}: {len(unused)} checkpoint tensors " + f"unused, e.g. {unused[:3]}." + ) + del state + gc.collect() text_encoder = HFCLIPTextModel.from_pretrained( repo, subfolder="text_encoder", **load ) text_state = { k if k.startswith("text_model.") else f"text_model.{k}": v - for k, v in numpy_state_dict(text_encoder).items() + for k, v in { + k: v.detach().cpu().numpy() for k, v in text_encoder.state_dict().items() + }.items() } CLIPTextModel.transfer_from_hf(model.text_encoder, text_state) del text_encoder @@ -254,7 +202,7 @@ def build_from_diffusers(repo, token=None): for variant, source in sources.items(): print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") - model, config = build_from_diffusers(source, token=token) + model, config = transfer_stable_diffusion(source, token=token) n_bytes = sum(int(np.prod(w.shape)) * 4 for w in model.weights) stem = os.path.join(OUT_DIR, variant.replace("-", "_")) diff --git a/zeromodels/models/stable_diffusion/stable_diffusion_config.py b/zeromodels/models/stable_diffusion/stable_diffusion_config.py index 79a96b8c..ad5c6128 100644 --- a/zeromodels/models/stable_diffusion/stable_diffusion_config.py +++ b/zeromodels/models/stable_diffusion/stable_diffusion_config.py @@ -7,7 +7,9 @@ class UNet2DConditionConfig(BaseConfig): The defaults match the Stable Diffusion 1.x UNet (860M parameters). Fields mirror the model constructor and serialize flat; build a model from it with - ``UNet2DConditionModel(config)``. + ``UNet2DConditionModel(config)``. The same class serves later UNets that only + change the widths (Stable Diffusion 2.x: 1024-d cross-attention, one head count + per level, a linear token projection). Args: sample_size (`int`, *optional*, defaults to 64): @@ -23,10 +25,29 @@ class UNet2DConditionConfig(BaseConfig): ResNet blocks per down level (up levels use one more). cross_attention_dim (`int`, *optional*, defaults to 768): Width of the text ``encoder_hidden_states``. - num_attention_heads (`int`, *optional*, defaults to 8): - Attention heads in the ``CrossAttn`` blocks. + num_attention_heads (`int` or `tuple`, *optional*, defaults to 8): + Attention heads in the ``CrossAttn`` blocks, one value for every level + or a tuple with one per level. norm_num_groups (`int`, *optional*, defaults to 32): GroupNorm group count. + use_linear_projection (`bool`, *optional*, defaults to False): + Project the Transformer2D tokens with a linear layer instead of a 1x1 + convolution on the feature map. + transformer_layers_per_block (`int` or `tuple`, *optional*, defaults to 1): + Transformer blocks stacked in each Transformer2D, one value for every + level or one per level (SDXL: (1, 2, 10)). + addition_embed_type (`str`, *optional*): + `"text_time"` adds SDXL's micro-conditioning to the timestep embedding + (the pooled text embedding plus the sinusoidally embedded size / crop + `time_ids`, through the `add_embedding` MLP); `None` for SD 1.x / 2.x. + addition_time_embed_dim (`int`, *optional*, defaults to 256): + Sinusoidal embedding width of each time id. + projection_class_embeddings_input_dim (`int`, *optional*): + Input width of the `add_embedding` MLP: the pooled text width plus + `num_time_ids * addition_time_embed_dim` (SDXL: 1280 + 6 * 256 = 2816). + num_time_ids (`int`, *optional*, defaults to 6): + Micro-conditioning values per image (original size, crop offset, target + size). text_seq_len (`int`, *optional*, defaults to 77): Static text sequence length (CLIP pads to 77). @@ -61,8 +82,14 @@ class UNet2DConditionConfig(BaseConfig): block_out_channels: tuple = (320, 640, 1280, 1280) layers_per_block: int = 2 cross_attention_dim: int = 768 - num_attention_heads: int = 8 + num_attention_heads: int | tuple = 8 norm_num_groups: int = 32 + use_linear_projection: bool = False + transformer_layers_per_block: int | tuple = 1 + addition_embed_type: str | None = None + addition_time_embed_dim: int = 256 + projection_class_embeddings_input_dim: int | None = None + num_time_ids: int = 6 text_seq_len: int = 77 @@ -87,6 +114,13 @@ class AutoencoderKLConfig(BaseConfig): Image resolution the encoder/decoder graphs are built for. scaling_factor (`float`, *optional*, defaults to 0.18215): Latent scaling applied by the pipeline around the VAE. + force_upcast (`bool`, *optional*, defaults to False): + Build the VAE in float32 whatever dtype the rest of the model loads in + (the SDXL VAE overflows in float16). + shift_factor (`float`, *optional*, defaults to 0.0): + Latent offset applied with the scaling (SD3: ``(z - shift) * scale``). + use_quant_conv / use_post_quant_conv (`bool`, *optional*, defaults to True): + The 1x1 convolutions around the latent (absent in the SD3 VAE). Examples: @@ -109,6 +143,10 @@ class AutoencoderKLConfig(BaseConfig): norm_num_groups: int = 32 sample_size: int = 512 scaling_factor: float = 0.18215 + force_upcast: bool = False + shift_factor: float = 0.0 + use_quant_conv: bool = True + use_post_quant_conv: bool = True class StableDiffusionTextConfig(CLIPTextConfig): diff --git a/zeromodels/models/stable_diffusion/stable_diffusion_layers.py b/zeromodels/models/stable_diffusion/stable_diffusion_layers.py index 1ff117db..63ba4f8c 100644 --- a/zeromodels/models/stable_diffusion/stable_diffusion_layers.py +++ b/zeromodels/models/stable_diffusion/stable_diffusion_layers.py @@ -51,7 +51,7 @@ def group_norm(x, name, channels_axis, groups=GROUPS, eps=1e-5): @keras.saving.register_keras_serializable(package="zeromodels") -class ResnetBlock2D(layers.Layer): +class StableDiffusionResnetBlock2D(layers.Layer): """Diffusers ``ResnetBlock2D``: GroupNorm/SiLU/conv, additive time embedding, GroupNorm/SiLU/conv, plus a 1x1 shortcut when the channel count changes. @@ -187,7 +187,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class CrossAttention(layers.Layer): +class StableDiffusionCrossAttention(layers.Layer): """Diffusers ``Attention`` (``to_q`` / ``to_k`` / ``to_v`` / ``to_out.0``) over ``(B, N, C)`` tokens; self-attention when called without a context. @@ -258,7 +258,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class GEGLUFeedForward(layers.Layer): +class StableDiffusionGEGLUFeedForward(layers.Layer): """Diffusers ``FeedForward`` with GEGLU: ``proj`` to ``2 * inner`` gated by GELU, then back down (``net.0.proj`` / ``net.2``). @@ -302,7 +302,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class BasicTransformerBlock(layers.Layer): +class StableDiffusionBasicTransformerBlock(layers.Layer): """Self-attention, cross-attention, GEGLU feed-forward, each pre-normed with a residual (diffusers ``BasicTransformerBlock``). @@ -321,15 +321,19 @@ def __init__(self, dim, heads, module_path, **kwargs): self.norm1 = layers.LayerNormalization( epsilon=1e-5, name=safe_name(f"{module_path}.norm1") ) - self.attn1 = CrossAttention(dim, heads, module_path=f"{module_path}.attn1") + self.attn1 = StableDiffusionCrossAttention( + dim, heads, module_path=f"{module_path}.attn1" + ) self.norm2 = layers.LayerNormalization( epsilon=1e-5, name=safe_name(f"{module_path}.norm2") ) - self.attn2 = CrossAttention(dim, heads, module_path=f"{module_path}.attn2") + self.attn2 = StableDiffusionCrossAttention( + dim, heads, module_path=f"{module_path}.attn2" + ) self.norm3 = layers.LayerNormalization( epsilon=1e-5, name=safe_name(f"{module_path}.norm3") ) - self.ff = GEGLUFeedForward(dim, module_path=f"{module_path}.ff") + self.ff = StableDiffusionGEGLUFeedForward(dim, module_path=f"{module_path}.ff") def build(self, input_shape): x_shape, context_shape = input_shape @@ -359,9 +363,9 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class Transformer2DModel(layers.Layer): - """Diffusers ``Transformer2DModel``: GroupNorm, 1x1 in-proj, one - ``BasicTransformerBlock`` over the flattened spatial tokens, 1x1 out-proj, +class StableDiffusionTransformer2DModel(layers.Layer): + """Diffusers ``Transformer2DModel``: GroupNorm, in-proj, ``num_layers`` + ``StableDiffusionBasicTransformerBlock`` over the flattened spatial tokens, out-proj, residual. ``call([x, context])`` with ``x`` in the active image data format. Args: @@ -369,6 +373,10 @@ class Transformer2DModel(layers.Layer): heads: Attention heads. module_path: Diffusers module path (``down_blocks.0.attentions.0``). groups: GroupNorm groups. + use_linear_projection: Project the flattened tokens with a linear layer + instead of the feature map with a 1x1 conv. + num_layers: Transformer blocks in the stack (1 for SD 1.x / 2.x; SDXL + stacks up to 10). data_format: ``"channels_last"`` or ``"channels_first"``; defaults to ``keras.config.image_data_format()``. channels_axis: The channel axis of that layout (``-1`` or ``1``). @@ -380,6 +388,8 @@ def __init__( heads, module_path, groups=GROUPS, + use_linear_projection=False, + num_layers=1, data_format=None, channels_axis=None, **kwargs, @@ -390,6 +400,8 @@ def __init__( self.heads = heads self.module_path = module_path self.groups = groups + self.use_linear_projection = use_linear_projection + self.num_layers = num_layers self.data_format = data_format or keras.config.image_data_format() self.channels_axis = ( channels_axis @@ -402,23 +414,44 @@ def __init__( epsilon=GROUP_EPS, name=safe_name(f"{module_path}.norm"), ) - self.proj_in = layers.Conv2D( - channels, - 1, - padding="valid", - data_format=self.data_format, - name=safe_name(f"{module_path}.proj_in"), - ) - self.transformer_block = BasicTransformerBlock( - channels, heads, module_path=f"{module_path}.transformer_blocks.0" - ) - self.proj_out = layers.Conv2D( - channels, - 1, - padding="valid", - data_format=self.data_format, - name=safe_name(f"{module_path}.proj_out"), - ) + if use_linear_projection: + self.proj_in = layers.Dense( + channels, name=safe_name(f"{module_path}.proj_in") + ) + self.proj_out = layers.Dense( + channels, name=safe_name(f"{module_path}.proj_out") + ) + else: + self.proj_in = layers.Conv2D( + channels, + 1, + padding="valid", + data_format=self.data_format, + name=safe_name(f"{module_path}.proj_in"), + ) + self.proj_out = layers.Conv2D( + channels, + 1, + padding="valid", + data_format=self.data_format, + name=safe_name(f"{module_path}.proj_out"), + ) + # Block 0 keeps the single-block attribute name (the hosted SD 1.x / 2.x + # checkpoints' h5 layout follows attribute names); the deeper SDXL stacks + # add transformer_block_1, transformer_block_2, ... + for k in range(num_layers): + block = StableDiffusionBasicTransformerBlock( + channels, heads, module_path=f"{module_path}.transformer_blocks.{k}" + ) + setattr( + self, "transformer_block" if k == 0 else f"transformer_block_{k}", block + ) + + @property + def transformer_blocks(self): + return [self.transformer_block] + [ + getattr(self, f"transformer_block_{k}") for k in range(1, self.num_layers) + ] def build(self, input_shape): x_shape, context_shape = input_shape @@ -428,17 +461,24 @@ def build(self, input_shape): else: height, width = x_shape[1], x_shape[2] proj_shape = (x_shape[0], height, width, self.channels) + tokens_shape = (x_shape[0], height * width, self.channels) self.norm.build(x_shape) - self.proj_in.build(x_shape) - self.transformer_block.build( - ((x_shape[0], height * width, self.channels), context_shape) - ) - self.proj_out.build(proj_shape) + if self.use_linear_projection: + # the projections act on the (B, H*W, C) tokens + self.proj_in.build(tokens_shape) + self.proj_out.build(tokens_shape) + else: + self.proj_in.build(x_shape) + self.proj_out.build(proj_shape) + for block in self.transformer_blocks: + block.build((tokens_shape, context_shape)) self.built = True def call(self, inputs): x, context = inputs - h = self.proj_in(self.norm(x)) + h = self.norm(x) + if not self.use_linear_projection: + h = self.proj_in(h) # 1x1 conv on the feature map shape = ops.shape(h) if self.data_format == "channels_first": # (B, C, H, W) -> (B, H*W, C) @@ -447,11 +487,18 @@ def call(self, inputs): else: height, width = shape[1], shape[2] h = ops.reshape(h, (-1, height * width, self.channels)) - h = self.transformer_block([h, context]) + if self.use_linear_projection: + h = self.proj_in(h) # linear on the tokens + for block in self.transformer_blocks: + h = block([h, context]) + if self.use_linear_projection: + h = self.proj_out(h) h = ops.reshape(h, (-1, height, width, self.channels)) if self.data_format == "channels_first": h = ops.transpose(h, (0, 3, 1, 2)) - return x + self.proj_out(h) + if not self.use_linear_projection: + h = self.proj_out(h) + return x + h def compute_output_shape(self, input_shape): return tuple(input_shape[0]) @@ -464,6 +511,8 @@ def get_config(self): "heads": self.heads, "module_path": self.module_path, "groups": self.groups, + "use_linear_projection": self.use_linear_projection, + "num_layers": self.num_layers, "data_format": self.data_format, "channels_axis": self.channels_axis, } @@ -472,7 +521,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class Downsample2D(layers.Layer): +class StableDiffusionDownsample2D(layers.Layer): """Strided 3x3 conv downsample (diffusers ``Downsample2D``). Args: @@ -560,7 +609,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class Upsample2D(layers.Layer): +class StableDiffusionUpsample2D(layers.Layer): """Nearest 2x upsample + 3x3 conv (diffusers ``Upsample2D``). Args: @@ -628,7 +677,7 @@ def get_config(self): @keras.saving.register_keras_serializable(package="zeromodels") -class VaeAttentionBlock(layers.Layer): +class StableDiffusionVaeAttentionBlock(layers.Layer): """Single-head spatial self-attention of the VAE mid block: GroupNorm, attention over the flattened pixels (biased q/k/v), residual. @@ -667,7 +716,7 @@ def __init__( epsilon=GROUP_EPS, name=safe_name(f"{module_path}.group_norm"), ) - self.attention = CrossAttention( + self.attention = StableDiffusionCrossAttention( channels, heads=1, module_path=module_path, diff --git a/zeromodels/models/stable_diffusion/stable_diffusion_model.py b/zeromodels/models/stable_diffusion/stable_diffusion_model.py index 5d449ee5..7956b43e 100644 --- a/zeromodels/models/stable_diffusion/stable_diffusion_model.py +++ b/zeromodels/models/stable_diffusion/stable_diffusion_model.py @@ -3,13 +3,10 @@ from zeromodels.base import BaseModel from zeromodels.base.base_diffusion import BaseDiffusion +from zeromodels.base.base_mixin import build_dtype_scope from zeromodels.base.base_scheduler import PNDMScheduler, get_scheduler from zeromodels.models.clip import CLIPTextModel -from .convert_stable_diffusion_diffusers_to_keras import ( - RenamedStateDict, - transfer_component, -) from .stable_diffusion_config import ( AutoencoderKLConfig, StableDiffusionConfig, @@ -17,11 +14,11 @@ ) from .stable_diffusion_layers import ( GROUP_EPS, - Downsample2D, - ResnetBlock2D, - Transformer2DModel, - Upsample2D, - VaeAttentionBlock, + StableDiffusionDownsample2D, + StableDiffusionResnetBlock2D, + StableDiffusionTransformer2DModel, + StableDiffusionUpsample2D, + StableDiffusionVaeAttentionBlock, group_norm, safe_name, time_embedding_mlp, @@ -50,16 +47,25 @@ class UNet2DConditionModel(BaseModel): latents to the text ``encoder_hidden_states``. It predicts the noise (or ``v``) on a latent, and is called once per denoising step by the pipeline. - Inputs are a dict ``{"sample", "timestep", "encoder_hidden_states"}``; the - output dict is ``{"sample": (B, H, W, out_channels)}``. Built for a fixed latent - resolution (``sample_size``), since the Transformer2D reshapes need static - spatial dims; the weights are resolution-independent, so rebuild the graph for - another size and reload. + Inputs are a dict ``{"sample", "timestep", "encoder_hidden_states"}`` (plus + ``"text_embeds"`` and ``"time_ids"`` for the SDXL ``text_time`` conditioning); + the output dict is ``{"sample": (B, H, W, out_channels)}``. Built for a fixed + latent resolution (``sample_size``), since the Transformer2D reshapes need + static spatial dims; the weights are resolution-independent, so rebuild the + graph for another size and reload. Args mirror the diffusers config: ``in_channels`` / ``out_channels`` (4), ``block_out_channels`` ((320, 640, 1280, 1280)), ``layers_per_block`` (2), - ``cross_attention_dim`` (768), ``num_attention_heads`` (8), ``norm_num_groups`` - (32), and the ``down_block_types`` / ``up_block_types`` lists. + ``cross_attention_dim`` (768), ``num_attention_heads`` (8, or one value per + level), ``norm_num_groups`` (32), ``use_linear_projection`` (a linear token + projection instead of the 1x1 conv), ``transformer_layers_per_block`` (1, or + one value per level: SDXL stacks (1, 2, 10) blocks per Transformer2D) and the + ``down_block_types`` / ``up_block_types`` lists. ``addition_embed_type`` + ``"text_time"`` adds SDXL's micro-conditioning to the timestep embedding: the + pooled text embedding and the ``num_time_ids`` size / crop values, each + sinusoidally embedded to ``addition_time_embed_dim``, concatenated + (``projection_class_embeddings_input_dim`` wide) and projected by the + ``add_embedding`` MLP. """ HF_MODEL_TYPE = None @@ -77,6 +83,12 @@ def __init__( cross_attention_dim=768, num_attention_heads=8, norm_num_groups=32, + use_linear_projection=False, + transformer_layers_per_block=1, + addition_embed_type=None, + addition_time_embed_dim=256, + projection_class_embeddings_input_dim=None, + num_time_ids=6, text_seq_len=77, data_format=None, channels_axis=None, @@ -86,7 +98,16 @@ def __init__( data_format = data_format or keras.config.image_data_format() if channels_axis is None: channels_axis = -1 if data_format == "channels_last" else 1 - heads = num_attention_heads + levels = len(block_out_channels) + # heads per level (diffusers' attention_head_dim list), mirrored on the way up + if isinstance(num_attention_heads, (tuple, list)): + heads_per_level = tuple(num_attention_heads) + else: + heads_per_level = (num_attention_heads,) * levels + if isinstance(transformer_layers_per_block, (tuple, list)): + depth_per_level = tuple(transformer_layers_per_block) + else: + depth_per_level = (transformer_layers_per_block,) * levels time_embed_dim = block_out_channels[0] * 4 sample_h, sample_w = ( sample_size @@ -108,6 +129,33 @@ def __init__( temb = timestep_embedding(timestep_in, block_out_channels[0]) temb = time_embedding_mlp(temb, time_embed_dim, name="time_embedding") + extra_inputs = {} + if addition_embed_type == "text_time": + # SDXL micro-conditioning: the pooled text embedding and the sinusoidal + # embeddings of the size / crop ids, projected and added to temb + pooled_dim = ( + projection_class_embeddings_input_dim + - num_time_ids * addition_time_embed_dim + ) + text_embeds_in = layers.Input(shape=(pooled_dim,), name="text_embeds") + time_ids_in = layers.Input(shape=(num_time_ids,), name="time_ids") + time_embeds = timestep_embedding( + ops.reshape(time_ids_in, (-1,)), addition_time_embed_dim + ) + time_embeds = ops.reshape( + time_embeds, (-1, num_time_ids * addition_time_embed_dim) + ) + add_embeds = ops.concatenate([text_embeds_in, time_embeds], axis=-1) + temb = temb + time_embedding_mlp( + add_embeds, time_embed_dim, name="add_embedding" + ) + extra_inputs = {"text_embeds": text_embeds_in, "time_ids": time_ids_in} + elif addition_embed_type is not None: + raise ValueError( + f"Unsupported addition_embed_type {addition_embed_type!r}; expected " + "None or 'text_time'." + ) + sample = layers.Conv2D( block_out_channels[0], 3, @@ -120,7 +168,7 @@ def __init__( for i, block_type in enumerate(down_block_types): out_ch = block_out_channels[i] for j in range(layers_per_block): - sample = ResnetBlock2D( + sample = StableDiffusionResnetBlock2D( out_ch, module_path=f"down_blocks.{i}.resnets.{j}", groups=norm_num_groups, @@ -128,17 +176,19 @@ def __init__( channels_axis=channels_axis, )([sample, temb]) if block_type == CROSS_ATTN_DOWN: - sample = Transformer2DModel( + sample = StableDiffusionTransformer2DModel( out_ch, - heads, + heads_per_level[i], module_path=f"down_blocks.{i}.attentions.{j}", groups=norm_num_groups, + use_linear_projection=use_linear_projection, + num_layers=depth_per_level[i], data_format=data_format, channels_axis=channels_axis, )([sample, context]) skips.append(sample) if i != len(down_block_types) - 1: - sample = Downsample2D( + sample = StableDiffusionDownsample2D( out_ch, module_path=f"down_blocks.{i}.downsamplers.0", data_format=data_format, @@ -147,22 +197,24 @@ def __init__( skips.append(sample) mid_ch = block_out_channels[-1] - sample = ResnetBlock2D( + sample = StableDiffusionResnetBlock2D( mid_ch, module_path="mid_block.resnets.0", groups=norm_num_groups, data_format=data_format, channels_axis=channels_axis, )([sample, temb]) - sample = Transformer2DModel( + sample = StableDiffusionTransformer2DModel( mid_ch, - heads, + heads_per_level[-1], module_path="mid_block.attentions.0", groups=norm_num_groups, + use_linear_projection=use_linear_projection, + num_layers=depth_per_level[-1], data_format=data_format, channels_axis=channels_axis, )([sample, context]) - sample = ResnetBlock2D( + sample = StableDiffusionResnetBlock2D( mid_ch, module_path="mid_block.resnets.1", groups=norm_num_groups, @@ -171,11 +223,13 @@ def __init__( )([sample, temb]) reversed_channels = list(reversed(block_out_channels)) + reversed_heads = list(reversed(heads_per_level)) + reversed_depths = list(reversed(depth_per_level)) for i, block_type in enumerate(up_block_types): out_ch = reversed_channels[i] for j in range(layers_per_block + 1): sample = layers.Concatenate(axis=channels_axis)([sample, skips.pop()]) - sample = ResnetBlock2D( + sample = StableDiffusionResnetBlock2D( out_ch, module_path=f"up_blocks.{i}.resnets.{j}", groups=norm_num_groups, @@ -183,16 +237,18 @@ def __init__( channels_axis=channels_axis, )([sample, temb]) if block_type == CROSS_ATTN_UP: - sample = Transformer2DModel( + sample = StableDiffusionTransformer2DModel( out_ch, - heads, + reversed_heads[i], module_path=f"up_blocks.{i}.attentions.{j}", groups=norm_num_groups, + use_linear_projection=use_linear_projection, + num_layers=reversed_depths[i], data_format=data_format, channels_axis=channels_axis, )([sample, context]) if i != len(up_block_types) - 1: - sample = Upsample2D( + sample = StableDiffusionUpsample2D( out_ch, module_path=f"up_blocks.{i}.upsamplers.0", data_format=data_format, @@ -219,6 +275,7 @@ def __init__( "sample": sample_in, "timestep": timestep_in, "encoder_hidden_states": context, + **extra_inputs, }, outputs={"sample": sample}, name=name, @@ -235,33 +292,65 @@ def __init__( self.block_out_channels = tuple(block_out_channels) self.layers_per_block = layers_per_block self.cross_attention_dim = cross_attention_dim - self.num_attention_heads = num_attention_heads + self.num_attention_heads = ( + tuple(num_attention_heads) + if isinstance(num_attention_heads, (tuple, list)) + else num_attention_heads + ) self.norm_num_groups = norm_num_groups + self.use_linear_projection = use_linear_projection + self.transformer_layers_per_block = ( + tuple(transformer_layers_per_block) + if isinstance(transformer_layers_per_block, (tuple, list)) + else transformer_layers_per_block + ) + self.addition_embed_type = addition_embed_type + self.addition_time_embed_dim = addition_time_embed_dim + self.projection_class_embeddings_input_dim = ( + projection_class_embeddings_input_dim + ) + self.num_time_ids = num_time_ids self.text_seq_len = text_seq_len @classmethod def from_diffusers_config(cls, config, **kwargs): """Build from a diffusers ``unet/config.json`` dict.""" + return cls(**cls.kwargs_from_diffusers_config(config), **kwargs) + + @staticmethod + def kwargs_from_diffusers_config(config): + """The constructor kwargs a diffusers ``unet/config.json`` dict describes.""" + # diffusers' "attention_head_dim" is the head COUNT (legacy name), a scalar + # or one value per level heads = config.get("num_attention_heads") or config.get("attention_head_dim", 8) if isinstance(heads, (list, tuple)): - heads = heads[0] - return cls( - sample_size=config.get("sample_size", 64), - in_channels=config.get("in_channels", 4), - out_channels=config.get("out_channels", 4), - down_block_types=tuple(config["down_block_types"]), - up_block_types=tuple(config["up_block_types"]), - block_out_channels=tuple(config["block_out_channels"]), - layers_per_block=config.get("layers_per_block", 2), - cross_attention_dim=config.get("cross_attention_dim", 768), - num_attention_heads=heads, - norm_num_groups=config.get("norm_num_groups", 32), - **kwargs, - ) - - @classmethod - def transfer_from_hf(cls, keras_model, state_dict): - transfer_component(keras_model, state_dict) + heads = tuple(heads) + depth = config.get("transformer_layers_per_block", 1) + if isinstance(depth, (list, tuple)): + depth = tuple(depth) + kwargs = { + "sample_size": config.get("sample_size", 64), + "in_channels": config.get("in_channels", 4), + "out_channels": config.get("out_channels", 4), + "down_block_types": tuple(config["down_block_types"]), + "up_block_types": tuple(config["up_block_types"]), + "block_out_channels": tuple(config["block_out_channels"]), + "layers_per_block": config.get("layers_per_block", 2), + "cross_attention_dim": config.get("cross_attention_dim", 768), + "num_attention_heads": heads, + "norm_num_groups": config.get("norm_num_groups", 32), + "use_linear_projection": config.get("use_linear_projection", False), + "transformer_layers_per_block": depth, + } + if config.get("addition_embed_type") is not None: + kwargs.update( + addition_embed_type=config["addition_embed_type"], + addition_time_embed_dim=config.get("addition_time_embed_dim", 256), + projection_class_embeddings_input_dim=config[ + "projection_class_embeddings_input_dim" + ], + ) + return kwargs def get_config(self): config = super().get_config() @@ -277,6 +366,14 @@ def get_config(self): "cross_attention_dim": self.cross_attention_dim, "num_attention_heads": self.num_attention_heads, "norm_num_groups": self.norm_num_groups, + "use_linear_projection": self.use_linear_projection, + "transformer_layers_per_block": self.transformer_layers_per_block, + "addition_embed_type": self.addition_embed_type, + "addition_time_embed_dim": self.addition_time_embed_dim, + "projection_class_embeddings_input_dim": ( + self.projection_class_embeddings_input_dim + ), + "num_time_ids": self.num_time_ids, "text_seq_len": self.text_seq_len, "name": self.name, } @@ -285,8 +382,8 @@ def get_config(self): def vae_resnet(channels, module_path, groups, data_format, channels_axis): - """A VAE ``ResnetBlock2D``: no timestep conditioning, GroupNorm eps 1e-6.""" - return ResnetBlock2D( + """A VAE ``StableDiffusionResnetBlock2D``: no timestep conditioning, GroupNorm eps 1e-6.""" + return StableDiffusionResnetBlock2D( channels, module_path=module_path, groups=groups, @@ -339,7 +436,7 @@ def build_vae_encoder( )(x) if i != len(block_out_channels) - 1: # the VAE pads bottom/right only before its stride-2 conv (padding=0) - x = Downsample2D( + x = StableDiffusionDownsample2D( ch, module_path=f"encoder.down_blocks.{i}.downsamplers.0", padding=0, @@ -350,7 +447,7 @@ def build_vae_encoder( x = vae_resnet( mid, "encoder.mid_block.resnets.0", groups, data_format, channels_axis )(x) - x = VaeAttentionBlock( + x = StableDiffusionVaeAttentionBlock( mid, module_path="encoder.mid_block.attentions.0", groups=groups, @@ -413,7 +510,7 @@ def build_vae_decoder( x = vae_resnet( mid, "decoder.mid_block.resnets.0", groups, data_format, channels_axis )(x) - x = VaeAttentionBlock( + x = StableDiffusionVaeAttentionBlock( mid, module_path="decoder.mid_block.attentions.0", groups=groups, @@ -434,7 +531,7 @@ def build_vae_decoder( channels_axis, )(x) if i != len(reversed_channels) - 1: - x = Upsample2D( + x = StableDiffusionUpsample2D( ch, module_path=f"decoder.up_blocks.{i}.upsamplers.0", data_format=data_format, @@ -479,8 +576,12 @@ class AutoencoderKL(BaseModel): Args mirror the diffusers config: ``in_channels`` / ``out_channels`` (3), ``latent_channels`` (4), ``block_out_channels`` ((128, 256, 512, 512)), ``layers_per_block`` (2), ``norm_num_groups`` (32), ``sample_size`` (the image - resolution the encoder/decoder graphs are built for), and ``scaling_factor`` - (0.18215). + resolution the encoder/decoder graphs are built for), ``scaling_factor`` + (0.18215), ``force_upcast`` (build this VAE in float32 whatever dtype the + rest of the model loads in: the SDXL VAE overflows in float16), + ``shift_factor`` (the SD3 latent offset) and ``use_quant_conv`` / + ``use_post_quant_conv`` (the 1x1 convolutions around the latent, absent in the + SD3 VAE). """ config_class = AutoencoderKLConfig @@ -496,6 +597,10 @@ def __init__( norm_num_groups=32, sample_size=512, scaling_factor=0.18215, + force_upcast=False, + shift_factor=0.0, + use_quant_conv=True, + use_post_quant_conv=True, data_format=None, channels_axis=None, name="AutoencoderKL", @@ -512,60 +617,75 @@ def __init__( factor = 2 ** (len(block_out_channels) - 1) h_lat, w_lat = h_img // factor, w_img // factor - encoder = build_vae_encoder( - (h_img, w_img), - in_channels, - block_out_channels, - layers_per_block, - latent_channels, - norm_num_groups, - data_format, - channels_axis, - ) - decoder = build_vae_decoder( - (h_lat, w_lat), - out_channels, - block_out_channels, - layers_per_block, - latent_channels, - norm_num_groups, - data_format, - channels_axis, - ) - quant_conv = layers.Conv2D( - 2 * latent_channels, - 1, - data_format=data_format, - name="quant_conv", - ) - post_quant_conv = layers.Conv2D( - latent_channels, - 1, - data_format=data_format, - name="post_quant_conv", - ) + # force_upcast: the layers are created under a float32 policy whatever the + # global (load) dtype is, as diffusers upcasts such a VAE around decode + with build_dtype_scope("float32" if force_upcast else None): + encoder = build_vae_encoder( + (h_img, w_img), + in_channels, + block_out_channels, + layers_per_block, + latent_channels, + norm_num_groups, + data_format, + channels_axis, + ) + decoder = build_vae_decoder( + (h_lat, w_lat), + out_channels, + block_out_channels, + layers_per_block, + latent_channels, + norm_num_groups, + data_format, + channels_axis, + ) + quant_conv = ( + layers.Conv2D( + 2 * latent_channels, + 1, + data_format=data_format, + name="quant_conv", + ) + if use_quant_conv + else None + ) + post_quant_conv = ( + layers.Conv2D( + latent_channels, + 1, + data_format=data_format, + name="post_quant_conv", + ) + if use_post_quant_conv + else None + ) - image_in = layers.Input( - shape=(in_channels, h_img, w_img) - if data_format == "channels_first" - else (h_img, w_img, in_channels), - name="image", - ) - latent_in = layers.Input( - shape=(latent_channels, h_lat, w_lat) - if data_format == "channels_first" - else (h_lat, w_lat, latent_channels), - name="latent", - ) - moments = quant_conv(encoder(image_in)) - decoded = decoder(post_quant_conv(latent_in)) + image_in = layers.Input( + shape=(in_channels, h_img, w_img) + if data_format == "channels_first" + else (h_img, w_img, in_channels), + name="image", + ) + latent_in = layers.Input( + shape=(latent_channels, h_lat, w_lat) + if data_format == "channels_first" + else (h_lat, w_lat, latent_channels), + name="latent", + ) + moments = encoder(image_in) + if quant_conv is not None: + moments = quant_conv(moments) + decoded = decoder( + latent_in if post_quant_conv is None else post_quant_conv(latent_in) + ) - super().__init__( - inputs={"image": image_in, "latent": latent_in}, - outputs={"moments": moments, "sample": decoded}, - name=name, - **kwargs, - ) + super().__init__( + inputs={"image": image_in, "latent": latent_in}, + outputs={"moments": moments, "sample": decoded}, + name=name, + **kwargs, + ) self.data_format = data_format self.channels_axis = channels_axis @@ -577,6 +697,10 @@ def __init__( self.norm_num_groups = norm_num_groups self.sample_size = sample_size self.scaling_factor = scaling_factor + self.force_upcast = force_upcast + self.shift_factor = shift_factor + self.use_quant_conv = use_quant_conv + self.use_post_quant_conv = use_post_quant_conv self.vae_scale_factor = factor self.encoder = encoder self.decoder = decoder @@ -584,7 +708,9 @@ def __init__( self.post_quant_conv = post_quant_conv def encode(self, image, sample=False, seed=None): - moments = self.quant_conv(self.encoder(image)) + moments = self.encoder(image) + if self.quant_conv is not None: + moments = self.quant_conv(moments) mean, logvar = ops.split(moments, 2, axis=self.channels_axis) if not sample: return mean @@ -594,7 +720,9 @@ def encode(self, image, sample=False, seed=None): return mean + std * noise def decode(self, latent): - return self.decoder(self.post_quant_conv(latent)) + if self.post_quant_conv is not None: + latent = self.post_quant_conv(latent) + return self.decoder(latent) def get_config(self): config = super().get_config() @@ -612,13 +740,13 @@ def from_diffusers_config(cls, config, sample_size=512, **kwargs): norm_num_groups=config.get("norm_num_groups", 32), sample_size=sample_size, scaling_factor=config.get("scaling_factor", 0.18215), + force_upcast=config.get("force_upcast", False), + shift_factor=config.get("shift_factor") or 0.0, + use_quant_conv=config.get("use_quant_conv", True), + use_post_quant_conv=config.get("use_post_quant_conv", True), **kwargs, ) - def transfer_from_hf(self, state_dict): - # legacy query/key/value/proj_attn spelling of a raw checkpoint read - transfer_component(self, RenamedStateDict(state_dict)) - @keras.saving.register_keras_serializable(package="zeromodels") class StableDiffusionModel(BaseModel): @@ -657,12 +785,24 @@ class StableDiffusionModel(BaseModel): def __init__(self, name="StableDiffusionModel", **kwargs): keras_kwargs = {k: kwargs.pop(k) for k in ("trainable", "dtype") if k in kwargs} - config = StableDiffusionConfig.from_dict(kwargs) # regroup the flat kwargs - u, v, t = config.unet_config, config.vae_config, config.text_config + config = self.config_class.from_dict(kwargs) # regroup the flat kwargs data_format = keras.config.image_data_format() channels_axis = -1 if data_format == "channels_last" else 1 + components = self.build_components(config, data_format, channels_axis) + inputs, outputs = self.build_graph(config, components, data_format) + super().__init__(inputs=inputs, outputs=outputs, name=name, **keras_kwargs) + + self.data_format = data_format + self.channels_axis = channels_axis + for attr, component in components.items(): + setattr(self, attr, component) + + def build_components(self, config, data_format, channels_axis): + """The trained components, ``{attribute name: model}``: the UNet, the VAE + and the CLIP text tower (subclasses swap or add towers).""" + u, v, t = config.unet_config, config.vae_config, config.text_config unet = UNet2DConditionModel( u, data_format=data_format, channels_axis=channels_axis ) @@ -681,7 +821,12 @@ def __init__(self, name="StableDiffusionModel", **kwargs): hidden_act=config.hidden_act, layer_norm_eps=config.layer_norm_eps, ) + return {"unet": unet, "vae": vae, "text_encoder": text_encoder} + def unet_vae_inputs(self, config, components, data_format): + """The UNet's and the VAE's ``keras.Input`` dict (shared by every SD family).""" + u, v = config.unet_config, config.vae_config + vae = components["vae"] sample_h, sample_w = ( u.sample_size if isinstance(u.sample_size, (tuple, list)) @@ -693,69 +838,68 @@ def __init__(self, name="StableDiffusionModel", **kwargs): else (v.sample_size, v.sample_size) ) lat_h, lat_w = img_h // vae.vae_scale_factor, img_w // vae.vae_scale_factor - - sample_in = layers.Input( - shape=(u.in_channels, sample_h, sample_w) - if data_format == "channels_first" - else (sample_h, sample_w, u.in_channels), - name="sample", + inputs = { + "sample": layers.Input( + shape=(u.in_channels, sample_h, sample_w) + if data_format == "channels_first" + else (sample_h, sample_w, u.in_channels), + name="sample", + ), + "timestep": layers.Input(shape=(), name="timestep"), + "encoder_hidden_states": layers.Input( + shape=(u.text_seq_len, u.cross_attention_dim), + name="encoder_hidden_states", + ), + "image": layers.Input( + shape=(v.in_channels, img_h, img_w) + if data_format == "channels_first" + else (img_h, img_w, v.in_channels), + name="image", + ), + "latent": layers.Input( + shape=(v.latent_channels, lat_h, lat_w) + if data_format == "channels_first" + else (lat_h, lat_w, v.latent_channels), + name="latent", + ), + } + return inputs + + def build_graph(self, config, components, data_format): + """Wire the components into the container's ``(inputs, outputs)`` dicts: + three disconnected paths, one per component.""" + t = config.text_config + unet, vae, text_encoder = ( + components["unet"], + components["vae"], + components["text_encoder"], ) - timestep_in = layers.Input(shape=(), name="timestep") - context_in = layers.Input( - shape=(u.text_seq_len, u.cross_attention_dim), name="encoder_hidden_states" + inputs = self.unet_vae_inputs(config, components, data_format) + inputs["token_ids"] = layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="token_ids" ) - image_in = layers.Input( - shape=(v.in_channels, img_h, img_w) - if data_format == "channels_first" - else (img_h, img_w, v.in_channels), - name="image", + inputs["padding_mask"] = layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="padding_mask" ) - latent_in = layers.Input( - shape=(v.latent_channels, lat_h, lat_w) - if data_format == "channels_first" - else (lat_h, lat_w, v.latent_channels), - name="latent", - ) - token_ids_in = layers.Input(shape=(t.max_seq_len,), name="token_ids") - padding_mask_in = layers.Input(shape=(t.max_seq_len,), name="padding_mask") noise_pred = unet( { - "sample": sample_in, - "timestep": timestep_in, - "encoder_hidden_states": context_in, + "sample": inputs["sample"], + "timestep": inputs["timestep"], + "encoder_hidden_states": inputs["encoder_hidden_states"], } )["sample"] - vae_out = vae({"image": image_in, "latent": latent_in}) + vae_out = vae({"image": inputs["image"], "latent": inputs["latent"]}) text_embeds = text_encoder( - {"token_ids": token_ids_in, "padding_mask": padding_mask_in} + {"token_ids": inputs["token_ids"], "padding_mask": inputs["padding_mask"]} )["last_hidden_state"] - - super().__init__( - inputs={ - "sample": sample_in, - "timestep": timestep_in, - "encoder_hidden_states": context_in, - "image": image_in, - "latent": latent_in, - "token_ids": token_ids_in, - "padding_mask": padding_mask_in, - }, - outputs={ - "noise_pred": noise_pred, - "moments": vae_out["moments"], - "image": vae_out["sample"], - "text_embeds": text_embeds, - }, - name=name, - **keras_kwargs, - ) - - self.data_format = data_format - self.channels_axis = channels_axis - self.unet = unet - self.vae = vae - self.text_encoder = text_encoder + outputs = { + "noise_pred": noise_pred, + "moments": vae_out["moments"], + "image": vae_out["sample"], + "text_embeds": text_embeds, + } + return inputs, outputs @classmethod def from_hf(cls, repo, **kwargs): @@ -827,10 +971,14 @@ def __init__(self, scheduler=None, name="StableDiffusionTextToImage", **kwargs): scheduler = ( get_scheduler(scheduler_config) if scheduler_config - else PNDMScheduler(steps_offset=1) + else self.default_scheduler() ) self.scheduler = scheduler + def default_scheduler(self): + """The family's sampler when the config carries no ``scheduler_config``.""" + return PNDMScheduler(steps_offset=1) + @property def latent_shape(self): size = self.unet.sample_size @@ -863,7 +1011,19 @@ def predict_noise(self, latents, timesteps, embeddings): )["sample"] def decode_latents(self, latents): - image = self.vae.decode(latents / self.vae.scaling_factor) + image = self.vae.decode( + latents / self.vae.scaling_factor + self.vae.shift_factor + ) if self.data_format == "channels_first": image = ops.transpose(image, (0, 2, 3, 1)) # generate() hands out HWC return image + + def encode_latents(self, image): + # the posterior mean, scaled: the image-to-image starting point (the + # reference samples the posterior; its noise is negligible next to the + # noise added for the strength) + if self.data_format == "channels_first": + image = ops.transpose(image, (0, 3, 1, 2)) # generate() hands in HWC + return ( + self.vae.encode(image) - self.vae.shift_factor + ) * self.vae.scaling_factor diff --git a/zeromodels/models/stable_diffusion/stable_diffusion_tokenizer.py b/zeromodels/models/stable_diffusion/stable_diffusion_tokenizer.py index f16868d3..6a20d1a0 100644 --- a/zeromodels/models/stable_diffusion/stable_diffusion_tokenizer.py +++ b/zeromodels/models/stable_diffusion/stable_diffusion_tokenizer.py @@ -19,9 +19,17 @@ class StableDiffusionTokenizer(CLIPTokenizer): ``tokenizer_file`` is given (no default repo). tokenizer_file: Explicit ``tokenizer.json`` path (overrides ``hf_id``). max_seq_len: Padded / truncated length (default 77). + pad_token: Pad token string (``<|endoftext|>``). """ - def __init__(self, hf_id=None, tokenizer_file=None, max_seq_len=77, **kwargs): + def __init__( + self, + hf_id=None, + tokenizer_file=None, + max_seq_len=77, + pad_token="<|endoftext|>", + **kwargs, + ): if tokenizer_file is None and hf_id is None: raise ValueError( f"{type(self).__name__}() needs an hf_id (a hosted repo, read from its " @@ -32,7 +40,10 @@ def __init__(self, hf_id=None, tokenizer_file=None, max_seq_len=77, **kwargs): if tokenizer_file is None: tokenizer_file = self.download_tokenizer_json(hf_id) super().__init__( - tokenizer_file=tokenizer_file, max_seq_len=max_seq_len, **kwargs + tokenizer_file=tokenizer_file, + max_seq_len=max_seq_len, + pad_token=pad_token, + **kwargs, ) self.hf_id = hf_id diff --git a/zeromodels/models/stable_diffusion_2/__init__.py b/zeromodels/models/stable_diffusion_2/__init__.py new file mode 100644 index 00000000..45b64e47 --- /dev/null +++ b/zeromodels/models/stable_diffusion_2/__init__.py @@ -0,0 +1,19 @@ +from .stable_diffusion_2_config import ( + StableDiffusion2Config, + StableDiffusion2TextConfig, + StableDiffusion2UNetConfig, +) +from .stable_diffusion_2_model import ( + StableDiffusion2Model, + StableDiffusion2TextToImage, +) +from .stable_diffusion_2_tokenizer import StableDiffusion2Tokenizer + +__all__ = [ + "StableDiffusion2Model", + "StableDiffusion2TextToImage", + "StableDiffusion2Config", + "StableDiffusion2TextConfig", + "StableDiffusion2UNetConfig", + "StableDiffusion2Tokenizer", +] diff --git a/zeromodels/models/stable_diffusion_2/convert_stable_diffusion_2_diffusers_to_keras.py b/zeromodels/models/stable_diffusion_2/convert_stable_diffusion_2_diffusers_to_keras.py new file mode 100644 index 00000000..97ff389b --- /dev/null +++ b/zeromodels/models/stable_diffusion_2/convert_stable_diffusion_2_diffusers_to_keras.py @@ -0,0 +1,87 @@ +import numpy as np + +from zeromodels.models.stable_diffusion.convert_stable_diffusion_diffusers_to_keras import ( + config_from_diffusers as stable_diffusion_config_from_diffusers, +) +from zeromodels.models.stable_diffusion.convert_stable_diffusion_diffusers_to_keras import ( + transfer_stable_diffusion, +) + +# Stable Diffusion 2.x: one architecture, five checkpoints (the 768px ones are +# v-prediction models built at sample_size 96; SD-Turbo is the 512px 2.1 distilled +# to a few steps). The stabilityai SD 2 repos are no longer on the Hub; +# sd2-community mirrors them in the same diffusers layout. This script only +# converts weights, the SD 1.x way (see that converter); the hosting tooling +# writes zm_config.json / tokenizer.json. +STABLE_DIFFUSION_2_SOURCES = { + "stable-diffusion-2-base": "sd2-community/stable-diffusion-2-base", + "stable-diffusion-2": "sd2-community/stable-diffusion-2", + "stable-diffusion-2-1-base": "sd2-community/stable-diffusion-2-1-base", + "stable-diffusion-2-1": "sd2-community/stable-diffusion-2-1", + # SD 2.1 distilled for 1 to 4 steps without guidance (Euler, trailing spacing) + "sd-turbo": "stabilityai/sd-turbo", +} + + +def config_from_diffusers(repo, token=None): + from zeromodels.models.stable_diffusion_2.stable_diffusion_2_config import ( + StableDiffusion2Config, + ) + + return stable_diffusion_config_from_diffusers( + repo, token=token, config_cls=StableDiffusion2Config + ) + + +def transfer_stable_diffusion_2(repo, token=None): + from zeromodels.models.stable_diffusion_2.stable_diffusion_2_config import ( + StableDiffusion2Config, + ) + from zeromodels.models.stable_diffusion_2.stable_diffusion_2_model import ( + StableDiffusion2Model, + ) + + return transfer_stable_diffusion( + repo, + token=token, + model_cls=StableDiffusion2Model, + config_cls=StableDiffusion2Config, + ) + + +if __name__ == "__main__": + import gc + import os + + import keras + + OUT_DIR = os.environ.get("ZM_OUT_DIR", "C:/Users/gites/Desktop/code/v1_weights") + os.makedirs(OUT_DIR, exist_ok=True) + MAX_SHARD_GB = 5.0 + token = os.environ.get("HF_TOKEN") + # ZM_VARIANTS="stable-diffusion-2-1,stable-diffusion-2-1-base" converts a subset + # (each variant is a ~5 GB download and a ~5 GB output); default: all four. + selected = [v for v in os.environ.get("ZM_VARIANTS", "").split(",") if v] + sources = { + variant: source + for variant, source in STABLE_DIFFUSION_2_SOURCES.items() + if not selected or variant in selected + } + + for variant, source in sources.items(): + print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") + model, config = transfer_stable_diffusion_2(source, token=token) + + n_bytes = sum(int(np.prod(w.shape)) * 4 for w in model.weights) + stem = os.path.join(OUT_DIR, variant.replace("-", "_")) + if n_bytes > MAX_SHARD_GB * 1024**3: + out = f"{stem}.weights.json" + model.save_weights(out, max_shard_size=MAX_SHARD_GB) + else: + out = f"{stem}.weights.h5" + model.save_weights(out) + print(f" saved -> {out} ({n_bytes / 1024**3:.2f} GB fp32)") + + del model + keras.backend.clear_session() + gc.collect() diff --git a/zeromodels/models/stable_diffusion_2/stable_diffusion_2_config.py b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_config.py new file mode 100644 index 00000000..edf18cba --- /dev/null +++ b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_config.py @@ -0,0 +1,79 @@ +from zeromodels.models.clip.clip_config import CLIPTextConfig +from zeromodels.models.stable_diffusion.stable_diffusion_config import ( + StableDiffusionConfig, + UNet2DConditionConfig, +) + + +class StableDiffusion2UNetConfig(UNet2DConditionConfig): + r"""The Stable Diffusion 2.x denoiser: a [`UNet2DConditionConfig`] with the SD 2 + widths. + + Same fields as [`UNet2DConditionConfig`]; the defaults are the SD 2.x UNet (865M + parameters): a 1024-d text context (OpenCLIP ViT-H/14), one head count per level + ((5, 10, 20, 20), a 64-wide head everywhere) and a linear token projection in the + Transformer2D blocks. ``sample_size`` is 64 for the 512px checkpoints and 96 for + the 768px ones.""" + + cross_attention_dim: int = 1024 + num_attention_heads: int | tuple = (5, 10, 20, 20) + use_linear_projection: bool = True + + +class StableDiffusion2TextConfig(CLIPTextConfig): + r"""The Stable Diffusion 2.x text tower: OpenCLIP ViT-H/14's text encoder. + + Same fields as [`CLIPTextConfig`], defaulted to the ViT-H/14 sizes (1024 wide, + 16 heads) with 23 layers: the checkpoints ship the encoder truncated to its + penultimate layer, which SD 2 conditions on. The activation is ``gelu`` (set + through the top-level ``hidden_act``).""" + + hidden_dim: int = 1024 + num_heads: int = 16 + num_layers: int = 23 + + +class StableDiffusion2Config(StableDiffusionConfig): + r"""Configuration for [`StableDiffusion2Model`], the hosted SD 2.x weights container. + + Same shape as [`StableDiffusionConfig`] (one nested ``unet_config`` / + ``vae_config`` / ``text_config`` plus the scheduler and token ids; flat + constructor), with the SD 2.x defaults: the [`StableDiffusion2UNetConfig`] + denoiser, the [`StableDiffusion2TextConfig`] OpenCLIP-H text tower with ``gelu``, + and ``!`` (id 0) as the pad token. The four SD 2.x checkpoints (2-base, 2, + 2-1-base, 2-1) share this architecture up to ``sample_size``; the 768px ones ship + a v-prediction scheduler. + + Args: + unet_config ([`StableDiffusion2UNetConfig`], *optional*): The denoiser. + vae_config ([`AutoencoderKLConfig`], *optional*): The VAE. + text_config ([`StableDiffusion2TextConfig`], *optional*): The OpenCLIP + ViT-H/14 text encoder. + hidden_act (`str`, *optional*, defaults to `"gelu"`): + Text encoder activation. + pad_token_id (`int`, *optional*, defaults to 0): + OpenCLIP pads with ``!`` (id 0) rather than ``<|endoftext|>``. + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_2 import StableDiffusion2Config, StableDiffusion2Model + + >>> config = StableDiffusion2Config() + >>> model = StableDiffusion2Model(config) # random weights, SD 2.x shapes + >>> model.config.unet_config.num_attention_heads + (5, 10, 20, 20) + ```""" + + model_type = "stable_diffusion_2" + + sub_configs = { + "unet_config": StableDiffusion2UNetConfig, + "vae_config": StableDiffusionConfig.sub_configs["vae_config"], + "text_config": StableDiffusion2TextConfig, + } + + unet_config: StableDiffusion2UNetConfig | dict | None = None + text_config: StableDiffusion2TextConfig | dict | None = None + hidden_act: str = "gelu" + pad_token_id: int = 0 diff --git a/zeromodels/models/stable_diffusion_2/stable_diffusion_2_model.py b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_model.py new file mode 100644 index 00000000..e27ec14c --- /dev/null +++ b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_model.py @@ -0,0 +1,69 @@ +import keras + +from zeromodels.models.stable_diffusion.stable_diffusion_model import ( + StableDiffusionModel, + StableDiffusionTextToImage, +) + +from .stable_diffusion_2_config import StableDiffusion2Config + +# The container and the task build the identical graph, so both load the one +# hosted repo (whose zm_config.json names the container). +STABLE_DIFFUSION_2_HUB_SIBLINGS = frozenset( + {"StableDiffusion2Model", "StableDiffusion2TextToImage"} +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion2Model(StableDiffusionModel): + """Stable Diffusion 2.x weights container: the UNet, the VAE and the OpenCLIP + ViT-H/14 text encoder as one functional ``keras.Model``. + + The SD 1.x container with the SD 2.x configuration + (:class:`StableDiffusion2Config`): the same three disconnected sub-graphs and + the same components (``.unet`` / ``.vae`` / ``.text_encoder``), with a 1024-d + text context, one attention head count per UNet level, a linear token + projection in the Transformer2D blocks and the 23-layer ``gelu`` text tower. + Hosted as one ``zm_config.json`` + ``model.weights.h5`` per checkpoint under + ``zeromodels/stable-diffusion-2*``; on-the-fly ``hf:`` conversion is not + supported for diffusion models. + + Args: + **kwargs: The flat :class:`StableDiffusion2Config` fields (``unet_`` / + ``vae_`` / ``text_`` prefixed), or a config positionally. Pass + ``unet_sample_size=96, vae_sample_size=768`` for the 768px checkpoints + (their repos already carry it). + """ + + config_class = StableDiffusion2Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_2_HUB_SIBLINGS + + def __init__(self, name="StableDiffusion2Model", **kwargs): + super().__init__(name=name, **kwargs) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion2TextToImage(StableDiffusionTextToImage): + """Text-to-image Stable Diffusion 2.x, pure Keras 3 and cross-backend. + + :class:`StableDiffusionTextToImage` with the SD 2.x configuration: the same + ``generate`` (``BaseDiffusion``) over the SD 2 graph, so the checkpoint's + ``!``-padded empty prompt drives classifier-free guidance and its scheduler + (PNDM for ``stable-diffusion-2-1-base``, DDIM for the others, v-prediction for + the 768px checkpoints) comes from the repo's ``scheduler_config``:: + + sd = StableDiffusion2TextToImage.from_weights("zeromodels/stable-diffusion-2-1") + tokenizer = StableDiffusion2Tokenizer.from_weights("zeromodels/stable-diffusion-2-1") + image = sd.generate(**tokenizer("a photo of a cat")) # (1, 768, 768, 3) uint8 + + Args: + scheduler: A :class:`BaseScheduler`; defaults to the config's + ``scheduler_config``. + **kwargs: The flat :class:`StableDiffusion2Config` fields. + """ + + config_class = StableDiffusion2Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_2_HUB_SIBLINGS + + def __init__(self, scheduler=None, name="StableDiffusion2TextToImage", **kwargs): + super().__init__(scheduler=scheduler, name=name, **kwargs) diff --git a/zeromodels/models/stable_diffusion_2/stable_diffusion_2_tokenizer.py b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_tokenizer.py new file mode 100644 index 00000000..c9d33103 --- /dev/null +++ b/zeromodels/models/stable_diffusion_2/stable_diffusion_2_tokenizer.py @@ -0,0 +1,34 @@ +import keras + +from zeromodels.models.stable_diffusion.stable_diffusion_tokenizer import ( + StableDiffusionTokenizer, +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion2Tokenizer(StableDiffusionTokenizer): + """Stable Diffusion 2.x text tokenizer: the OpenCLIP ViT-H/14 BPE tokenizer. + + The same byte-pair encoding and ``<|startoftext|>`` / ``<|endoftext|>`` framing + as CLIP ViT-L/14's, truncated to 77 tokens, but padded with ``!`` (id 0) the + way OpenCLIP does, which is what the SD 2 text encoder was trained on. Loads by + repo id like the model: ``from_weights("zeromodels/stable-diffusion-2-1")``. + + Args: + hf_id: Hosted repo to read ``tokenizer.json`` from. Required unless + ``tokenizer_file`` is given (no default repo). + tokenizer_file: Explicit ``tokenizer.json`` path (overrides ``hf_id``). + max_seq_len: Padded / truncated length (default 77). + pad_token: Pad token string (``!``). + """ + + def __init__( + self, hf_id=None, tokenizer_file=None, max_seq_len=77, pad_token="!", **kwargs + ): + super().__init__( + hf_id=hf_id, + tokenizer_file=tokenizer_file, + max_seq_len=max_seq_len, + pad_token=pad_token, + **kwargs, + ) diff --git a/zeromodels/models/stable_diffusion_3/__init__.py b/zeromodels/models/stable_diffusion_3/__init__.py new file mode 100644 index 00000000..df0f4cbe --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/__init__.py @@ -0,0 +1,27 @@ +from .stable_diffusion_3_config import ( + StableDiffusion3Config, + StableDiffusion3T5EncoderConfig, + StableDiffusion3TextConfig, + StableDiffusion3TransformerConfig, + StableDiffusion3VAEConfig, +) +from .stable_diffusion_3_model import ( + SD3T5EncoderModel, + SD3Transformer2DModel, + StableDiffusion3Model, + StableDiffusion3TextToImage, +) +from .stable_diffusion_3_tokenizer import StableDiffusion3Tokenizer + +__all__ = [ + "SD3T5EncoderModel", + "SD3Transformer2DModel", + "StableDiffusion3Model", + "StableDiffusion3TextToImage", + "StableDiffusion3Config", + "StableDiffusion3T5EncoderConfig", + "StableDiffusion3TextConfig", + "StableDiffusion3TransformerConfig", + "StableDiffusion3VAEConfig", + "StableDiffusion3Tokenizer", +] diff --git a/zeromodels/models/stable_diffusion_3/convert_stable_diffusion_3_diffusers_to_keras.py b/zeromodels/models/stable_diffusion_3/convert_stable_diffusion_3_diffusers_to_keras.py new file mode 100644 index 00000000..73f0958d --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/convert_stable_diffusion_3_diffusers_to_keras.py @@ -0,0 +1,336 @@ +import re + +import numpy as np +from tqdm import tqdm + +from zeromodels.conversion.exceptions import ( + WeightMappingError, + WeightShapeMismatchError, +) +from zeromodels.conversion.weight_split_util import split_model_weights +from zeromodels.conversion.weight_transfer_util import ( + compare_keras_torch_names, + transfer_weights, +) +from zeromodels.models.stable_diffusion.convert_stable_diffusion_diffusers_to_keras import ( + WEIGHT_NAME_MAPPING, +) +from zeromodels.models.stable_diffusion_xl.convert_stable_diffusion_xl_diffusers_to_keras import ( + LazyNumpyStateDict, + text_config_kwargs, + transfer_text_encoder, +) + +STABLE_DIFFUSION_3_SOURCES = { + "stable-diffusion-3-medium": "stabilityai/stable-diffusion-3-medium-diffusers", +} +T5_XXL_ENCODER_VARIANT = "t5-v1_1-xxl-encoder" +T5_ENCODER_FIXED_LEAVES = { + "shared": "shared.weight", + "encoder_rel_bias": ( + "encoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight" + ), + "encoder_final_layer_norm": "encoder.final_layer_norm.weight", +} + + +def config_from_diffusers(repo, token=None, config_cls=None): + from diffusers import AutoencoderKL, FlowMatchEulerDiscreteScheduler + from diffusers import SD3Transformer2DModel as DiffusersSD3Transformer2DModel + from transformers import CLIPTextConfig, CLIPTokenizerFast, T5TokenizerFast + + from zeromodels.models.stable_diffusion_3.stable_diffusion_3_config import ( + StableDiffusion3Config, + ) + from zeromodels.models.stable_diffusion_3.stable_diffusion_3_model import ( + SD3Transformer2DModel, + ) + + config_cls = config_cls or StableDiffusion3Config + transformer = dict( + DiffusersSD3Transformer2DModel.load_config( + repo, subfolder="transformer", token=token + ) + ) + vae = dict(AutoencoderKL.load_config(repo, subfolder="vae", token=token)) + text = CLIPTextConfig.from_pretrained( + repo, subfolder="text_encoder", token=token + ).to_dict() + text_2 = CLIPTextConfig.from_pretrained( + repo, subfolder="text_encoder_2", token=token + ).to_dict() + scheduler = { + k: v + for k, v in FlowMatchEulerDiscreteScheduler.load_config( + repo, subfolder="scheduler", token=token + ).items() + if k == "_class_name" or not k.startswith("_") + } + tokenizer = CLIPTokenizerFast.from_pretrained( + repo, subfolder="tokenizer", token=token + ) + tokenizer_2 = CLIPTokenizerFast.from_pretrained( + repo, subfolder="tokenizer_2", token=token + ) + tokenizer_3 = T5TokenizerFast.from_pretrained( + repo, subfolder="tokenizer_3", token=token + ) + return config_cls( + transformer_config=SD3Transformer2DModel.kwargs_from_diffusers_config( + transformer + ), + vae_config={ + "in_channels": vae.get("in_channels", 3), + "out_channels": vae.get("out_channels", 3), + "latent_channels": vae.get("latent_channels", 16), + "block_out_channels": tuple(vae["block_out_channels"]), + "layers_per_block": vae.get("layers_per_block", 2), + "norm_num_groups": vae.get("norm_num_groups", 32), + "sample_size": transformer.get("sample_size", 128) + * 2 ** (len(vae["block_out_channels"]) - 1), + "scaling_factor": vae.get("scaling_factor") or 1.5305, + "shift_factor": vae.get("shift_factor") or 0.0, + "force_upcast": bool(vae.get("force_upcast", True)), + "use_quant_conv": bool(vae.get("use_quant_conv", True)), + "use_post_quant_conv": bool(vae.get("use_post_quant_conv", True)), + }, + text_config={ + **text_config_kwargs(text), + "projection_dim": text.get("projection_dim", 768), + "hidden_act": text.get("hidden_act", "quick_gelu"), + }, + text_config_2={ + **text_config_kwargs(text_2), + "projection_dim": text_2.get("projection_dim", 1280), + "hidden_act": text_2.get("hidden_act", "gelu"), + }, + hidden_act=text.get("hidden_act", "quick_gelu"), + layer_norm_eps=text.get("layer_norm_eps", 1e-5), + scheduler_config=scheduler, + bos_token_id=tokenizer.bos_token_id, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + pad_token_id_2=tokenizer_2.pad_token_id, + eos_token_id_3=tokenizer_3.eos_token_id, + pad_token_id_3=tokenizer_3.pad_token_id, + ) + + +def from_pretrained_fp16(cls, repo, subfolder, token=None): + import torch + + try: + return cls.from_pretrained( + repo, + subfolder=subfolder, + variant="fp16", + torch_dtype=torch.float16, + token=token, + ) + except OSError: + return cls.from_pretrained( + repo, subfolder=subfolder, torch_dtype=torch.float16, token=token + ) + + +def transfer_stable_diffusion_3( + repo, token=None, dtype="float16", model_cls=None, config_cls=None +): + import gc + + import torch + from diffusers import AutoencoderKL as DiffusersAutoencoderKL + from diffusers import SD3Transformer2DModel as DiffusersSD3Transformer2DModel + from transformers import CLIPTextModelWithProjection + + from zeromodels.base.base_mixin import build_dtype_scope + from zeromodels.models.stable_diffusion_3.stable_diffusion_3_model import ( + StableDiffusion3Model, + ) + + model_cls = model_cls or StableDiffusion3Model + config = config_from_diffusers(repo, token=token, config_cls=config_cls) + with build_dtype_scope(dtype): + model = model_cls(config) + + legacy = { + ".query.": ".to_q.", + ".key.": ".to_k.", + ".value.": ".to_v.", + ".proj_attn.": ".to_out.0.", + } + for component, subfolder in ( + (model.transformer, "transformer"), + (model.vae, "vae"), + ): + ignore = () + if subfolder == "transformer": + module = from_pretrained_fp16( + DiffusersSD3Transformer2DModel, repo, subfolder, token=token + ) + state = LazyNumpyStateDict(module) + table = component.get_layer("pos_embed").pos_embed + table.assign(np.asarray(state["pos_embed.pos_embed"]).reshape(table.shape)) + ignore = ("pos_embed.pos_embed",) + else: + module = DiffusersAutoencoderKL.from_pretrained( + repo, subfolder=subfolder, torch_dtype=torch.float32, token=token + ) + state = {} + for key, value in module.state_dict().items(): + if ".attentions." in key: + for old, new in legacy.items(): + key = key.replace(old, new) + state[key] = value.detach().cpu().numpy() + del module + consumed = set(ignore) + trainable, non_trainable = split_model_weights(component) + for keras_weight, _ in tqdm( + trainable + non_trainable, desc=f"Transferring {subfolder} weights to Keras" + ): + key = "/".join(keras_weight.path.split("/")[-2:]) + for old, new in WEIGHT_NAME_MAPPING.items(): + key = key.replace(old, new) + if key in ignore: + continue + consumed.add(key) + if key not in state: + raise WeightMappingError(keras_weight.path, key) + torch_weight = state[key] + if not compare_keras_torch_names( + keras_weight.path, keras_weight, key, torch_weight + ): + raise WeightShapeMismatchError( + keras_weight.path, keras_weight.shape, key, torch_weight.shape + ) + transfer_weights(key, keras_weight, torch_weight) + unused = sorted(set(state) - consumed) + if unused: + raise ValueError( + f"{type(component).__name__}: {len(unused)} checkpoint tensors " + f"unused, e.g. {unused[:3]}." + ) + del state + gc.collect() + + for attr, subfolder in ( + ("text_encoder", "text_encoder"), + ("text_encoder_2", "text_encoder_2"), + ): + text_encoder = CLIPTextModelWithProjection.from_pretrained( + repo, subfolder=subfolder, torch_dtype=torch.float32, token=token + ) + transfer_text_encoder(getattr(model, attr), text_encoder) + del text_encoder + gc.collect() + return model, config + + +def transfer_t5_encoder(repo, token=None, dtype="float16"): + import json + + from huggingface_hub import hf_hub_download + from safetensors import safe_open + from transformers import T5Config + + from zeromodels.base.base_mixin import build_dtype_scope + from zeromodels.models.stable_diffusion_3.stable_diffusion_3_model import ( + SD3T5EncoderModel, + ) + + hf_config = T5Config.from_pretrained(repo, subfolder="text_encoder_3", token=token) + cfg = SD3T5EncoderModel.kwargs_from_hf_config(hf_config.to_dict()) + with build_dtype_scope(dtype): + model = SD3T5EncoderModel(**cfg) + + from huggingface_hub.errors import EntryNotFoundError + + try: + index = hf_hub_download( + repo, + "model.safetensors.index.fp16.json", + subfolder="text_encoder_3", + token=token, + ) + except EntryNotFoundError: + index = hf_hub_download( + repo, + "model.safetensors.index.json", + subfolder="text_encoder_3", + token=token, + ) + with open(index, encoding="utf-8") as f: + weight_map = json.load(f)["weight_map"] + shard_paths = { + shard: hf_hub_download(repo, shard, subfolder="text_encoder_3", token=token) + for shard in sorted(set(weight_map.values())) + } + for weight in tqdm(model.weights, desc="Transferring weights to Keras"): + leaf = weight.path.split("/")[-2] + if leaf in T5_ENCODER_FIXED_LEAVES: + key = T5_ENCODER_FIXED_LEAVES[leaf] + else: + m = re.match(r"enc_(\d+)_(attn|ff)_(q|k|v|o|ln|wi_0|wi_1|wo)$", leaf) + if not m: + raise WeightMappingError(weight.path, leaf) + idx, role, part = m.groups() + sub, name = ( + ("0", "SelfAttention") if role == "attn" else ("1", "DenseReluDense") + ) + if part == "ln": + key = f"encoder.block.{idx}.layer.{sub}.layer_norm.weight" + else: + key = f"encoder.block.{idx}.layer.{sub}.{name}.{part}.weight" + with safe_open(shard_paths[weight_map[key]], framework="np") as shard: + transfer_weights(weight.path, weight, shard.get_tensor(key)) + return model, cfg + + +if __name__ == "__main__": + import gc + import os + + import keras + + OUT_DIR = os.environ.get("ZM_OUT_DIR", "C:/Users/gites/Desktop/code/v1_weights") + os.makedirs(OUT_DIR, exist_ok=True) + MAX_SHARD_GB = 5.0 + token = os.environ.get("HF_TOKEN") + selected = [v for v in os.environ.get("ZM_VARIANTS", "").split(",") if v] + sources = { + variant: source + for variant, source in STABLE_DIFFUSION_3_SOURCES.items() + if not selected or variant in selected + } + + def save(model, variant): + n_bytes = sum( + int(np.prod(w.shape)) * np.dtype(w.dtype).itemsize for w in model.weights + ) + stem = os.path.join(OUT_DIR, variant.replace("-", "_").replace(".", "_")) + if n_bytes > MAX_SHARD_GB * 1024**3: + out = f"{stem}.weights.json" + model.save_weights(out, max_shard_size=MAX_SHARD_GB) + else: + out = f"{stem}.weights.h5" + model.save_weights(out) + print(f" saved -> {out} ({n_bytes / 1024**3:.2f} GB)") + + for variant, source in sources.items(): + print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") + model, config = transfer_stable_diffusion_3(source, token=token) + save(model, variant) + del model + keras.backend.clear_session() + gc.collect() + + if not selected or T5_XXL_ENCODER_VARIANT in selected: + source = next(iter(STABLE_DIFFUSION_3_SOURCES.values())) + print( + f"\n{'=' * 60}\nConverting: {T5_XXL_ENCODER_VARIANT} <- {source}\n{'=' * 60}" + ) + model, cfg = transfer_t5_encoder(source, token=token) + save(model, T5_XXL_ENCODER_VARIANT) + del model + keras.backend.clear_session() + gc.collect() diff --git a/zeromodels/models/stable_diffusion_3/stable_diffusion_3_config.py b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_config.py new file mode 100644 index 00000000..3da5514f --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_config.py @@ -0,0 +1,236 @@ +from zeromodels.base import BaseConfig +from zeromodels.models.stable_diffusion.stable_diffusion_config import ( + AutoencoderKLConfig, +) +from zeromodels.models.stable_diffusion_xl.stable_diffusion_xl_config import ( + StableDiffusionXLTextConfig2, +) + + +class StableDiffusion3TransformerConfig(BaseConfig): + r"""Configuration for [`SD3Transformer2DModel`], the Stable Diffusion 3 MMDiT. + + The defaults match the SD 3 medium transformer (2B parameters). Fields mirror + the model constructor and serialize flat; build a model from it with + ``SD3Transformer2DModel(config)``. + + Args: + sample_size (`int`, *optional*, defaults to 128): + Latent spatial size the graph is built for (image size / 8). + patch_size (`int`, *optional*, defaults to 2): + Patch side; the token grid is ``sample_size / patch_size``. + in_channels / out_channels (`int`, *optional*, defaults to 16): + Latent channel counts in and out. + num_layers (`int`, *optional*, defaults to 24): + Joint transformer blocks. + attention_head_dim (`int`, *optional*, defaults to 64): + Width of one attention head. + num_attention_heads (`int`, *optional*, defaults to 24): + Attention heads; the token width is ``num_attention_heads * + attention_head_dim``. + joint_attention_dim (`int`, *optional*, defaults to 4096): + Width of the text features (the T5 width; the CLIP features are padded + to it). + caption_projection_dim (`int`, *optional*, defaults to 1536): + Width the text features are projected to (the token width). + pooled_projection_dim (`int`, *optional*, defaults to 2048): + Width of the pooled CLIP embeddings added to the timestep embedding. + pos_embed_max_size (`int`, *optional*, defaults to 192): + Side of the precomputed position grid the latent grid is cropped from. + qk_norm (`str`, *optional*): + ``"rms_norm"`` normalizes q and k per head (SD 3.5); `None` for SD 3. + dual_attention_layers (`tuple`, *optional*): + Blocks with a second, latent-only attention (SD 3.5 medium: 0 to 12). + text_seq_len (`int`, *optional*, defaults to 333): + Static text sequence length (77 CLIP + 256 T5 tokens). + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_3 import StableDiffusion3TransformerConfig, SD3Transformer2DModel + + >>> config = StableDiffusion3TransformerConfig() + >>> model = SD3Transformer2DModel(config) + >>> model.config.num_layers + 24 + ```""" + + model_type = "sd3_transformer_2d" + + sample_size: int = 128 + patch_size: int = 2 + in_channels: int = 16 + out_channels: int = 16 + num_layers: int = 24 + attention_head_dim: int = 64 + num_attention_heads: int = 24 + joint_attention_dim: int = 4096 + caption_projection_dim: int = 1536 + pooled_projection_dim: int = 2048 + pos_embed_max_size: int = 192 + qk_norm: str | None = None + dual_attention_layers: tuple = () + text_seq_len: int = 333 + + +class StableDiffusion3T5EncoderConfig(BaseConfig): + r"""Configuration for [`SD3T5EncoderModel`], the third Stable Diffusion 3 text + encoder: the encoder half of T5 v1.1 XXL. + + The defaults are the T5-XXL encoder every SD 3 / 3.5 checkpoint conditions on + (24 layers, 4096 wide, 64 heads of 64, gated-GELU feed-forward of 10240, 4.7B + parameters), hosted once as ``zeromodels/t5-v1_1-xxl-encoder``. Fields mirror + the model constructor and serialize flat. + + Args: + vocab_size (`int`, *optional*, defaults to 32128): + SentencePiece vocabulary. + embed_dim (`int`, *optional*, defaults to 4096): + Token width (``d_model``). + key_value_dim (`int`, *optional*, defaults to 64): + Width of one attention head (``d_kv``). + mlp_dim (`int`, *optional*, defaults to 10240): + Feed-forward inner width (``d_ff``). + num_layers (`int`, *optional*, defaults to 24): + Encoder blocks. + num_heads (`int`, *optional*, defaults to 64): + Attention heads. + relative_attention_num_buckets (`int`, *optional*, defaults to 32): + Relative-position-bias buckets. + relative_attention_max_distance (`int`, *optional*, defaults to 128): + Largest bucketed distance. + layer_norm_eps (`float`, *optional*, defaults to 1e-06): + RMSNorm epsilon. + pad_token_id / eos_token_id (`int`, *optional*, defaults to 0 / 1): + The tokenizer's ```` and ````. + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_3 import StableDiffusion3T5EncoderConfig, SD3T5EncoderModel + + >>> config = StableDiffusion3T5EncoderConfig(num_layers=2, embed_dim=64, mlp_dim=128, num_heads=4) + >>> model = SD3T5EncoderModel(config) + >>> model.config.num_layers + 2 + ```""" + + model_type = "stable_diffusion_3_t5_encoder" + + vocab_size: int = 32128 + embed_dim: int = 4096 + key_value_dim: int = 64 + mlp_dim: int = 10240 + num_layers: int = 24 + num_heads: int = 64 + relative_attention_num_buckets: int = 32 + relative_attention_max_distance: int = 128 + layer_norm_eps: float = 1e-06 + pad_token_id: int = 0 + eos_token_id: int = 1 + + +class StableDiffusion3VAEConfig(AutoencoderKLConfig): + r"""The Stable Diffusion 3 VAE: an [`AutoencoderKLConfig`] with 16 latent channels, + no quant convolutions, the SD 3 latent scaling (1.5305) and shift (0.0609), + built for 1024px in float32 (``force_upcast``).""" + + latent_channels: int = 16 + sample_size: int = 1024 + scaling_factor: float = 1.5305 + shift_factor: float = 0.0609 + force_upcast: bool = True + use_quant_conv: bool = False + use_post_quant_conv: bool = False + + +class StableDiffusion3TextConfig(StableDiffusionXLTextConfig2): + r"""The first Stable Diffusion 3 text tower: CLIP ViT-L/14's text encoder with + its 768-d projection (``CLIPTextModelWithProjection``), ``quick_gelu``.""" + + hidden_dim: int = 768 + num_heads: int = 12 + num_layers: int = 12 + projection_dim: int = 768 + hidden_act: str = "quick_gelu" + + +class StableDiffusion3Config(BaseConfig): + r"""Configuration for [`StableDiffusion3Model`], the hosted SD 3 weights container. + + One config for the container's components, so a repo carries a single + ``zm_config.json`` (nested ``transformer_config`` / ``vae_config`` / + ``text_config`` / ``text_config_2``) plus the weights, and loads with the + standard ``from_weights("zeromodels/")``. The constructor stays flat: + sub-config fields are prefixed ``transformer_`` / ``vae_`` / ``text_`` / + ``text_2_``. The third text encoder (T5-XXL) is not a component: it is a + separately hosted :class:`SD3T5EncoderModel` (its own + [`StableDiffusion3T5EncoderConfig`]) the task attaches on demand; + ``max_sequence_length`` and the T5 token ids describe its inputs. + + Args: + transformer_config ([`StableDiffusion3TransformerConfig`], *optional*): The + MMDiT. + vae_config ([`StableDiffusion3VAEConfig`], *optional*): The 16-channel VAE. + text_config ([`StableDiffusion3TextConfig`], *optional*): The CLIP ViT-L/14 + text encoder (with projection). + text_config_2 ([`StableDiffusionXLTextConfig2`], *optional*): The OpenCLIP + ViT-bigG/14 text encoder (with projection). + hidden_act (`str`, *optional*, defaults to `"quick_gelu"`): + First text tower activation (the second carries its own). + layer_norm_eps (`float`, *optional*, defaults to 1e-05): + Text encoder LayerNorm epsilon. + bos_token_id / eos_token_id / pad_token_id (`int`, *optional*): + CLIP's ``<|startoftext|>`` (49406) and ``<|endoftext|>`` (49407, also the + first tokenizer's pad). + pad_token_id_2 (`int`, *optional*, defaults to 0): + The second CLIP tokenizer pads with ``!`` (id 0). + max_sequence_length (`int`, *optional*, defaults to 256): + T5 token length. + eos_token_id_3 / pad_token_id_3 (`int`, *optional*, defaults to 1 / 0): + The T5 tokenizer's ```` and ````. + scheduler_config (`dict`, *optional*): + The checkpoint's sampler in the diffusers ``scheduler_config.json`` form + (``FlowMatchEulerDiscreteScheduler`` with its ``shift``). + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_3 import StableDiffusion3Config, StableDiffusion3Model + + >>> config = StableDiffusion3Config() + >>> model = StableDiffusion3Model(config) # random weights, SD 3 medium shapes + >>> model.config.transformer_config.num_layers + 24 + ```""" + + model_type = "stable_diffusion_3" + + sub_configs = { + "transformer_config": StableDiffusion3TransformerConfig, + "vae_config": StableDiffusion3VAEConfig, + "text_config": StableDiffusion3TextConfig, + "text_config_2": StableDiffusionXLTextConfig2, + } + sub_config_prefixes = { + "transformer_config": "transformer_", + "vae_config": "vae_", + "text_config": "text_", + "text_config_2": "text_2_", + } + group_extras = {"text_config": ("vocab_size", "max_seq_len")} + + transformer_config: StableDiffusion3TransformerConfig | dict | None = None + vae_config: StableDiffusion3VAEConfig | dict | None = None + text_config: StableDiffusion3TextConfig | dict | None = None + text_config_2: StableDiffusionXLTextConfig2 | dict | None = None + hidden_act: str = "quick_gelu" + layer_norm_eps: float = 1e-05 + bos_token_id: int = 49406 + eos_token_id: int = 49407 + pad_token_id: int = 49407 + pad_token_id_2: int = 0 + max_sequence_length: int = 256 + eos_token_id_3: int = 1 + pad_token_id_3: int = 0 + scheduler_config: dict | None = None diff --git a/zeromodels/models/stable_diffusion_3/stable_diffusion_3_layers.py b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_layers.py new file mode 100644 index 00000000..9b373c0d --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_layers.py @@ -0,0 +1,668 @@ +import keras +import numpy as np +from keras import layers, ops + +from zeromodels.base.base_attention import active_attn_implementation, fused_attention +from zeromodels.models.stable_diffusion.stable_diffusion_layers import safe_name +from zeromodels.models.t5.t5_layers import T5LayerNorm, T5SelfAttentionLayer + +NORM_EPS = 1e-6 + + +def sincos_pos_embed_2d(embed_dim, grid_size, base_size): + """The MMDiT 2D sinusoidal position table ``(grid_size**2, embed_dim)`` + (diffusers ``get_2d_sincos_pos_embed``): positions ``arange(grid_size) * + base_size / grid_size``, the first half of the channels embedding the column + and the second half the row, each as ``[sin, cos]`` over ``embed_dim // 4`` + geometric frequencies. Computed in float64 like the reference, stored float32. + """ + positions = np.arange(grid_size, dtype=np.float64) / (grid_size / base_size) + cols, rows = np.meshgrid(positions, positions, indexing="xy") # (H, W) each + half = embed_dim // 2 + omega = np.arange(half // 2, dtype=np.float64) / (half / 2.0) + omega = 1.0 / 10000**omega + + def embed(pos): + out = np.outer(pos.reshape(-1), omega) + return np.concatenate([np.sin(out), np.cos(out)], axis=1) + + table = np.concatenate([embed(cols), embed(rows)], axis=1) + return table.astype(np.float32) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3PatchEmbed(layers.Layer): + """MMDiT patch embedding (diffusers ``StableDiffusion3PatchEmbed`` with SD3 cropping): a + ``patch_size`` strided convolution flattens the latent to tokens, plus the + sinusoidal position table of a ``pos_embed_max_size`` grid cropped (centered) + to the actual latent grid. The table is a non-trainable weight: the checkpoints + carry their own (the released SD 3 tables are not the textbook grid, and the + reference loads them rather than recomputing), initialized to the sinusoid. + + Args: + embed_dim: Token width. + patch_size: Patch (and stride) size. + pos_embed_max_size: Side of the precomputed position grid. + base_size: The latent side (in patches) the table is scaled to. + module_path: Diffusers module path (``pos_embed``). + data_format / channels_axis: Layout of the latent. + """ + + def __init__( + self, + embed_dim, + patch_size, + pos_embed_max_size, + base_size, + module_path="pos_embed", + data_format=None, + channels_axis=None, + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.embed_dim = embed_dim + self.patch_size = patch_size + self.pos_embed_max_size = pos_embed_max_size + self.base_size = base_size + self.module_path = module_path + self.data_format = data_format or keras.config.image_data_format() + self.channels_axis = ( + channels_axis + if channels_axis is not None + else (-1 if self.data_format == "channels_last" else 1) + ) + self.proj = layers.Conv2D( + embed_dim, + patch_size, + strides=patch_size, + padding="valid", + data_format=self.data_format, + name=safe_name(f"{module_path}.proj"), + ) + + def grid(self, input_shape): + if self.data_format == "channels_first": + return input_shape[2] // self.patch_size, input_shape[3] // self.patch_size + return input_shape[1] // self.patch_size, input_shape[2] // self.patch_size + + def build(self, input_shape): + self.proj.build(input_shape) + self.height, self.width = self.grid(input_shape) + if max(self.height, self.width) > self.pos_embed_max_size: + raise ValueError( + f"A {self.height}x{self.width} latent grid exceeds the " + f"{self.pos_embed_max_size}-patch position table." + ) + self.pos_embed = self.add_weight( + name="pos_embed", + shape=(self.pos_embed_max_size**2, self.embed_dim), + initializer=keras.initializers.Constant( + sincos_pos_embed_2d( + self.embed_dim, self.pos_embed_max_size, self.base_size + ) + ), + trainable=False, + ) + self.built = True + + def call(self, x): + h = self.proj(x) + if self.data_format == "channels_first": + h = ops.transpose(h, (0, 2, 3, 1)) + h = ops.reshape(h, (-1, self.height * self.width, self.embed_dim)) + # the centered crop of the table (the graph is built for one grid) + top = (self.pos_embed_max_size - self.height) // 2 + left = (self.pos_embed_max_size - self.width) // 2 + table = ops.reshape( + self.pos_embed, + (self.pos_embed_max_size, self.pos_embed_max_size, self.embed_dim), + ) + table = table[top : top + self.height, left : left + self.width] + table = ops.reshape(table, (1, self.height * self.width, self.embed_dim)) + return h + ops.cast(table, h.dtype) + + def compute_output_shape(self, input_shape): + height, width = self.grid(input_shape) + return (input_shape[0], height * width, self.embed_dim) + + def get_config(self): + config = super().get_config() + config.update( + { + "embed_dim": self.embed_dim, + "patch_size": self.patch_size, + "pos_embed_max_size": self.pos_embed_max_size, + "base_size": self.base_size, + "module_path": self.module_path, + "data_format": self.data_format, + "channels_axis": self.channels_axis, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3AdaLayerNorm(layers.Layer): + """Adaptive LayerNorm of the MMDiT blocks: a LayerNorm without affine + parameters, modulated by ``num_chunks`` vectors projected from the SiLU'd + conditioning embedding (diffusers ``AdaLayerNormZero`` (6 chunks), + ``SD35AdaLayerNormZeroX`` (9, the dual-attention blocks) and + ``AdaLayerNormContinuous`` (2, the final norm and the last block's context)). + + ``call(x, temb)`` returns ``(x_modulated, gate_msa, shift_mlp, scale_mlp, + gate_mlp)`` for 6 chunks (plus ``x_modulated_2, gate_msa2`` for 9), or just the + modulated ``x`` for 2 chunks (``scale, shift`` order there). + + Args: + dim: Token width. + num_chunks: 2, 6 or 9. + module_path: Diffusers module path (``transformer_blocks.0.norm1``). + """ + + def __init__(self, dim, num_chunks, module_path, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.num_chunks = num_chunks + self.module_path = module_path + self.linear = layers.Dense( + num_chunks * dim, name=safe_name(f"{module_path}.linear") + ) + self.norm = layers.LayerNormalization( + epsilon=NORM_EPS, + center=False, + scale=False, + name=safe_name(f"{module_path}.norm"), + ) + + def build(self, input_shape): + x_shape, temb_shape = input_shape + self.linear.build(temb_shape) + self.norm.build(x_shape) + self.built = True + + def call(self, inputs): + x, temb = inputs + emb = self.linear(ops.silu(temb)) + chunks = ops.split(emb, self.num_chunks, axis=-1) + normed = self.norm(x) + if self.num_chunks == 2: + scale, shift = chunks + return normed * (1.0 + scale[:, None, :]) + shift[:, None, :] + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = chunks[:6] + modulated = normed * (1.0 + scale_msa[:, None, :]) + shift_msa[:, None, :] + outputs = [modulated, gate_msa, shift_mlp, scale_mlp, gate_mlp] + if self.num_chunks == 9: + shift_msa2, scale_msa2, gate_msa2 = chunks[6:] + outputs.append( + normed * (1.0 + scale_msa2[:, None, :]) + shift_msa2[:, None, :] + ) + outputs.append(gate_msa2) + return outputs + + def compute_output_shape(self, input_shape): + x_shape, temb_shape = input_shape + if self.num_chunks == 2: + return tuple(x_shape) + vector = (x_shape[0], self.dim) + outputs = [tuple(x_shape), vector, vector, vector, vector] + if self.num_chunks == 9: + outputs += [tuple(x_shape), vector] + return outputs + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "num_chunks": self.num_chunks, + "module_path": self.module_path, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3JointAttention(layers.Layer): + """MMDiT joint attention (diffusers ``Attention`` + ``JointAttnProcessor``): + the latent tokens and the text tokens get their own q / k / v projections + (``to_q`` / ``add_q_proj``, ...), attend jointly over the concatenated + sequence, and are projected back separately (``to_out.0`` / ``to_add_out``). + + Args: + dim: Token width (inner attention width). + heads: Attention heads. + module_path: Diffusers module path (``transformer_blocks.0.attn``). + context: Whether the text stream takes part (the dual-attention ``attn2`` + of SD 3.5 attends the latent tokens alone). + context_out: Whether the text stream is projected back (not in the last + block, whose text output is unused). + qk_norm: ``"rms_norm"`` normalizes q and k per head (SD 3.5), ``None`` not. + attn_implementation: The :func:`fused_attention` implementation used + when the model was built and is run without an explicit one + (``Model.from_weights(attn_implementation=...)``). ``"fused"``: the + joint sequence is 4429 tokens at 1024px, where the portable + ``"sdpa"`` math materializes ``(B, heads, 4429, 4429)`` float32 + logits (3.8 GB) per block. + """ + + def __init__( + self, + dim, + heads, + module_path, + context=True, + context_out=True, + qk_norm=None, + attn_implementation="fused", + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.heads = heads + self.module_path = module_path + self.context = context + self.context_out = context_out + self.qk_norm = qk_norm + self.attn_implementation = attn_implementation + self.head_dim = dim // heads + self.scale = self.head_dim**-0.5 + self.to_q, self.to_k, self.to_v = ( + self.dense("to_q"), + self.dense("to_k"), + self.dense("to_v"), + ) + self.to_out = self.dense("to_out.0") + if qk_norm == "rms_norm": + self.norm_q = self.rms_norm("norm_q") + self.norm_k = self.rms_norm("norm_k") + if context: + self.add_q_proj = self.dense("add_q_proj") + self.add_k_proj = self.dense("add_k_proj") + self.add_v_proj = self.dense("add_v_proj") + if qk_norm == "rms_norm": + self.norm_added_q = self.rms_norm("norm_added_q") + self.norm_added_k = self.rms_norm("norm_added_k") + if context_out: + self.to_add_out = self.dense("to_add_out") + + def dense(self, leaf): + return layers.Dense(self.dim, name=safe_name(f"{self.module_path}.{leaf}")) + + def rms_norm(self, leaf): + return layers.RMSNormalization( + epsilon=NORM_EPS, name=safe_name(f"{self.module_path}.{leaf}") + ) + + def build(self, input_shape, context_shape=None): + head_shape = (input_shape[0], self.heads, input_shape[1], self.head_dim) + for layer in (self.to_q, self.to_k, self.to_v, self.to_out): + layer.build(input_shape) + if self.qk_norm == "rms_norm": + self.norm_q.build(head_shape) + self.norm_k.build(head_shape) + if self.context: + ctx_heads = (context_shape[0], self.heads, context_shape[1], self.head_dim) + for layer in (self.add_q_proj, self.add_k_proj, self.add_v_proj): + layer.build(context_shape) + if self.qk_norm == "rms_norm": + self.norm_added_q.build(ctx_heads) + self.norm_added_k.build(ctx_heads) + if self.context_out: + self.to_add_out.build(context_shape) + self.built = True + + def split_heads(self, t): + t = ops.reshape(t, (-1, ops.shape(t)[1], self.heads, self.head_dim)) + return ops.transpose(t, (0, 2, 1, 3)) + + def call(self, x, context=None): + q = self.split_heads(self.to_q(x)) + k = self.split_heads(self.to_k(x)) + v = self.split_heads(self.to_v(x)) + if self.qk_norm == "rms_norm": + q, k = self.norm_q(q), self.norm_k(k) + n_x = ops.shape(x)[1] + if self.context: + cq = self.split_heads(self.add_q_proj(context)) + ck = self.split_heads(self.add_k_proj(context)) + cv = self.split_heads(self.add_v_proj(context)) + if self.qk_norm == "rms_norm": + cq, ck = self.norm_added_q(cq), self.norm_added_k(ck) + # the latent tokens first, then the text tokens, as the reference + q = ops.concatenate([q, cq], axis=2) + k = ops.concatenate([k, ck], axis=2) + v = ops.concatenate([v, cv], axis=2) + out = fused_attention( + q, + k, + v, + self.scale, + attn_implementation=active_attn_implementation() + or self.attn_implementation, + ) + out = ops.transpose(out, (0, 2, 1, 3)) + out = ops.reshape(out, (-1, ops.shape(out)[1], self.dim)) + if not self.context: + return self.to_out(out) + x_out = self.to_out(out[:, :n_x]) + if not self.context_out: + return x_out + return x_out, self.to_add_out(out[:, n_x:]) + + def compute_output_shape(self, input_shape, context_shape=None): + x_out = tuple(input_shape[:-1]) + (self.dim,) + if self.context and self.context_out: + return x_out, tuple(context_shape[:-1]) + (self.dim,) + return x_out + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "heads": self.heads, + "module_path": self.module_path, + "context": self.context, + "context_out": self.context_out, + "qk_norm": self.qk_norm, + "attn_implementation": self.attn_implementation, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3GELUFeedForward(layers.Layer): + """Diffusers ``FeedForward(activation_fn="gelu-approximate")``: ``net.0.proj`` + to ``mult * dim``, tanh GELU, ``net.2`` back to ``dim``. + + Args: + dim: Token width. + module_path: Diffusers module path (``transformer_blocks.0.ff``). + mult: Inner width multiplier (4). + """ + + def __init__(self, dim, module_path, mult=4, **kwargs): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.module_path = module_path + self.mult = mult + self.proj = layers.Dense( + dim * mult, name=safe_name(f"{module_path}.net.0.proj") + ) + self.out = layers.Dense(dim, name=safe_name(f"{module_path}.net.2")) + + def build(self, input_shape): + self.proj.build(input_shape) + self.out.build(tuple(input_shape[:-1]) + (self.dim * self.mult,)) + self.built = True + + def call(self, x): + return self.out(ops.gelu(self.proj(x), approximate=True)) + + def compute_output_shape(self, input_shape): + return tuple(input_shape[:-1]) + (self.dim,) + + def get_config(self): + config = super().get_config() + config.update( + {"dim": self.dim, "module_path": self.module_path, "mult": self.mult} + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3JointTransformerBlock(layers.Layer): + """One MMDiT block (diffusers ``StableDiffusion3JointTransformerBlock``): the latent and text + streams are each AdaLN-modulated by the conditioning embedding, attend jointly, + and pass through their own gated feed-forwards. ``call([x, context, temb])`` + returns ``[x, context]``, or ``x`` alone for the last block + (``context_pre_only``: its text stream is not updated). SD 3.5's dual-attention + blocks add a second, latent-only attention (``attn2``) with its own modulation. + + Args: + dim: Token width. + heads: Attention heads. + module_path: Diffusers module path (``transformer_blocks.0``). + context_pre_only: Last-block variant (2-chunk context norm, no text output). + qk_norm: ``"rms_norm"`` or ``None``. + use_dual_attention: The SD 3.5 dual-attention variant. + """ + + def __init__( + self, + dim, + heads, + module_path, + context_pre_only=False, + qk_norm=None, + use_dual_attention=False, + **kwargs, + ): + kwargs.setdefault("name", safe_name(module_path)) + super().__init__(**kwargs) + self.dim = dim + self.heads = heads + self.module_path = module_path + self.context_pre_only = context_pre_only + self.qk_norm = qk_norm + self.use_dual_attention = use_dual_attention + self.norm1 = StableDiffusion3AdaLayerNorm( + dim, 9 if use_dual_attention else 6, module_path=f"{module_path}.norm1" + ) + self.norm1_context = StableDiffusion3AdaLayerNorm( + dim, + 2 if context_pre_only else 6, + module_path=f"{module_path}.norm1_context", + ) + self.attn = StableDiffusion3JointAttention( + dim, + heads, + module_path=f"{module_path}.attn", + context=True, + context_out=not context_pre_only, + qk_norm=qk_norm, + ) + if use_dual_attention: + self.attn2 = StableDiffusion3JointAttention( + dim, + heads, + module_path=f"{module_path}.attn2", + context=False, + qk_norm=qk_norm, + ) + self.norm2 = layers.LayerNormalization( + epsilon=NORM_EPS, + center=False, + scale=False, + name=safe_name(f"{module_path}.norm2"), + ) + self.ff = StableDiffusion3GELUFeedForward(dim, module_path=f"{module_path}.ff") + if not context_pre_only: + self.norm2_context = layers.LayerNormalization( + epsilon=NORM_EPS, + center=False, + scale=False, + name=safe_name(f"{module_path}.norm2_context"), + ) + self.ff_context = StableDiffusion3GELUFeedForward( + dim, module_path=f"{module_path}.ff_context" + ) + + def build(self, input_shape): + x_shape, context_shape, temb_shape = input_shape + self.norm1.build((x_shape, temb_shape)) + self.norm1_context.build((context_shape, temb_shape)) + self.attn.build(x_shape, context_shape) + if self.use_dual_attention: + self.attn2.build(x_shape) + self.norm2.build(x_shape) + self.ff.build(x_shape) + if not self.context_pre_only: + self.norm2_context.build(context_shape) + self.ff_context.build(context_shape) + self.built = True + + def call(self, inputs): + x, context, temb = inputs + modulated = self.norm1([x, temb]) + normed_x, gate_msa, shift_mlp, scale_mlp, gate_mlp = modulated[:5] + if self.context_pre_only: + normed_context = self.norm1_context([context, temb]) + attn_x = self.attn(normed_x, normed_context) + else: + normed_context, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = ( + self.norm1_context([context, temb]) + ) + attn_x, attn_context = self.attn(normed_x, normed_context) + x = x + gate_msa[:, None, :] * attn_x + if self.use_dual_attention: + normed_x2, gate_msa2 = modulated[5:] + x = x + gate_msa2[:, None, :] * self.attn2(normed_x2) + h = self.norm2(x) * (1.0 + scale_mlp[:, None, :]) + shift_mlp[:, None, :] + x = x + gate_mlp[:, None, :] * self.ff(h) + if self.context_pre_only: + return x + context = context + c_gate_msa[:, None, :] * attn_context + h = self.norm2_context(context) * (1.0 + c_scale_mlp[:, None, :]) + h = h + c_shift_mlp[:, None, :] + context = context + c_gate_mlp[:, None, :] * self.ff_context(h) + return [x, context] + + def compute_output_shape(self, input_shape): + x_shape, context_shape, _ = input_shape + if self.context_pre_only: + return tuple(x_shape) + return [tuple(x_shape), tuple(context_shape)] + + def get_config(self): + config = super().get_config() + config.update( + { + "dim": self.dim, + "heads": self.heads, + "module_path": self.module_path, + "context_pre_only": self.context_pre_only, + "qk_norm": self.qk_norm, + "use_dual_attention": self.use_dual_attention, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3T5GatedFeedForward(layers.Layer): + """The T5 v1.1 feed-forward (``T5DenseGatedActDense``) with its pre-norm and + residual: ``x + wo(gelu_tanh(wi_0(norm(x))) * wi_1(norm(x)))``, bias-free + projections and the T5 RMSNorm. The original T5's ``relu(wi(x))`` block lives + in the t5 family; this gated form is what the SD 3 T5-XXL encoder uses. The + sub-layers are named ``{prefix}_ln`` / ``{prefix}_wi_0`` / ``{prefix}_wi_1`` / + ``{prefix}_wo`` so every weight path is unique across blocks. + + Args: + embed_dim: Token width. + mlp_dim: Inner width. + eps: RMSNorm epsilon. + prefix: Leaf-name prefix (``enc_3_ff``). + """ + + def __init__(self, embed_dim, mlp_dim, eps, prefix, **kwargs): + super().__init__(**kwargs) + self.embed_dim = embed_dim + self.mlp_dim = mlp_dim + self.eps = eps + self.prefix = prefix + self.layer_norm = T5LayerNorm(eps, name=f"{prefix}_ln") + self.wi_0 = layers.Dense(mlp_dim, use_bias=False, name=f"{prefix}_wi_0") + self.wi_1 = layers.Dense(mlp_dim, use_bias=False, name=f"{prefix}_wi_1") + self.wo = layers.Dense(embed_dim, use_bias=False, name=f"{prefix}_wo") + + def build(self, input_shape): + self.layer_norm.build(input_shape) + self.wi_0.build(input_shape) + self.wi_1.build(input_shape) + self.wo.build(tuple(input_shape[:-1]) + (self.mlp_dim,)) + self.built = True + + def call(self, hidden_states): + normed = self.layer_norm(hidden_states) + # "gelu_new": the tanh-approximate GELU of the T5 v1.1 checkpoints + inner = ops.gelu(self.wi_0(normed), approximate=True) * self.wi_1(normed) + return hidden_states + self.wo(inner) + + def compute_output_shape(self, input_shape): + return tuple(input_shape) + + def get_config(self): + config = super().get_config() + config.update( + { + "embed_dim": self.embed_dim, + "mlp_dim": self.mlp_dim, + "eps": self.eps, + "prefix": self.prefix, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3T5GatedEncoderBlock(layers.Layer): + """One T5 v1.1 encoder block: the t5 family's pre-norm self-attention with the + shared relative position bias, then :class:`StableDiffusion3T5GatedFeedForward`. + + Args: + embed_dim: Token width. + key_value_dim: Width of one attention head. + num_heads: Attention heads. + mlp_dim: Feed-forward inner width. + eps: RMSNorm epsilon. + prefix: Leaf-name prefix (``enc_3``). + """ + + def __init__( + self, embed_dim, key_value_dim, num_heads, mlp_dim, eps, prefix, **kwargs + ): + super().__init__(**kwargs) + self.embed_dim = embed_dim + self.key_value_dim = key_value_dim + self.num_heads = num_heads + self.mlp_dim = mlp_dim + self.eps = eps + self.prefix = prefix + self.self_attention = T5SelfAttentionLayer( + embed_dim, key_value_dim, num_heads, eps, prefix=f"{prefix}_attn" + ) + self.ff = StableDiffusion3T5GatedFeedForward( + embed_dim, mlp_dim, eps, prefix=f"{prefix}_ff" + ) + + def build(self, input_shape): + self.self_attention.build(input_shape) + self.ff.build(input_shape) + self.built = True + + def call(self, hidden_states, position_bias): + hidden_states = self.self_attention(hidden_states, position_bias) + return self.ff(hidden_states) + + def compute_output_spec(self, hidden_states, position_bias): + return keras.KerasTensor(hidden_states.shape, dtype=hidden_states.dtype) + + def get_config(self): + config = super().get_config() + config.update( + { + "embed_dim": self.embed_dim, + "key_value_dim": self.key_value_dim, + "num_heads": self.num_heads, + "mlp_dim": self.mlp_dim, + "eps": self.eps, + "prefix": self.prefix, + } + ) + return config diff --git a/zeromodels/models/stable_diffusion_3/stable_diffusion_3_model.py b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_model.py new file mode 100644 index 00000000..9ea4a7fe --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_model.py @@ -0,0 +1,720 @@ +import keras +from keras import layers, ops + +from zeromodels.base import BaseModel +from zeromodels.base.base_mixin import inference_scope +from zeromodels.base.base_scheduler import FlowMatchEulerDiscreteScheduler +from zeromodels.models.stable_diffusion.stable_diffusion_layers import ( + time_embedding_mlp, + timestep_embedding, +) +from zeromodels.models.stable_diffusion.stable_diffusion_model import ( + AutoencoderKL, + StableDiffusionModel, + StableDiffusionTextToImage, +) +from zeromodels.models.stable_diffusion_xl.stable_diffusion_xl_model import ( + build_text_encoder, +) +from zeromodels.models.t5.t5_layers import T5LayerNorm +from zeromodels.models.t5.t5_model import T5PositionBias + +from .stable_diffusion_3_config import ( + StableDiffusion3Config, + StableDiffusion3T5EncoderConfig, + StableDiffusion3TransformerConfig, +) +from .stable_diffusion_3_layers import ( + StableDiffusion3AdaLayerNorm, + StableDiffusion3JointTransformerBlock, + StableDiffusion3PatchEmbed, + StableDiffusion3T5GatedEncoderBlock, +) + +# The container and the task build the identical graph, so both load the one +# hosted repo (whose zm_config.json names the container). +STABLE_DIFFUSION_3_HUB_SIBLINGS = frozenset( + {"StableDiffusion3Model", "StableDiffusion3TextToImage"} +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class SD3Transformer2DModel(BaseModel): + """Stable Diffusion 3 denoiser, the MMDiT (diffusers ``SD3Transformer2DModel``). + + A multimodal diffusion transformer: the latent is patchified into tokens with + a cropped sinusoidal position table, the text tokens are projected to the same + width, and ``num_layers`` joint blocks let the two streams attend to each other + while a conditioning embedding (the sinusoidal timestep plus the pooled text + embedding) modulates every norm (AdaLN-Zero). The final norm and projection + unpatchify the latent tokens back to a velocity field of the latent's shape. + + Inputs are a dict ``{"sample", "timestep", "encoder_hidden_states", + "pooled_projections"}``; the output dict is ``{"sample": (B, H, W, C)}``. + Built for a fixed latent resolution (``sample_size``) like the SD UNet; the + weights are resolution-independent up to ``pos_embed_max_size`` patches. + + Args mirror the diffusers config: ``sample_size`` (128), ``patch_size`` (2), + ``in_channels`` / ``out_channels`` (16), ``num_layers`` (24), ``attention_head_dim`` + (64), ``num_attention_heads`` (24), ``joint_attention_dim`` (4096, the text + features' width), ``caption_projection_dim`` (the projected text width, the + token width), ``pooled_projection_dim`` (2048), ``pos_embed_max_size`` (192), + ``qk_norm`` (``"rms_norm"`` in SD 3.5), ``dual_attention_layers`` (the SD 3.5 + medium blocks with a second latent-only attention) and ``text_seq_len`` (the + static text length, 77 + 256). + """ + + HF_MODEL_TYPE = None + config_class = StableDiffusion3TransformerConfig + + def __init__( + self, + sample_size=128, + patch_size=2, + in_channels=16, + out_channels=16, + num_layers=24, + attention_head_dim=64, + num_attention_heads=24, + joint_attention_dim=4096, + caption_projection_dim=1536, + pooled_projection_dim=2048, + pos_embed_max_size=192, + qk_norm=None, + dual_attention_layers=(), + text_seq_len=333, + data_format=None, + channels_axis=None, + name="SD3Transformer2DModel", + **kwargs, + ): + data_format = data_format or keras.config.image_data_format() + if channels_axis is None: + channels_axis = -1 if data_format == "channels_last" else 1 + inner_dim = num_attention_heads * attention_head_dim + sample_h, sample_w = ( + sample_size + if isinstance(sample_size, (tuple, list)) + else (sample_size, sample_size) + ) + dual_attention_layers = tuple(dual_attention_layers) + + sample_in = layers.Input( + shape=(in_channels, sample_h, sample_w) + if data_format == "channels_first" + else (sample_h, sample_w, in_channels), + name="sample", + ) + timestep_in = layers.Input(shape=(), name="timestep") + context_in = layers.Input( + shape=(text_seq_len, joint_attention_dim), name="encoder_hidden_states" + ) + pooled_in = layers.Input( + shape=(pooled_projection_dim,), name="pooled_projections" + ) + + x = StableDiffusion3PatchEmbed( + inner_dim, + patch_size, + pos_embed_max_size, + base_size=sample_h // patch_size, + module_path="pos_embed", + data_format=data_format, + channels_axis=channels_axis, + )(sample_in) + temb = timestep_embedding(timestep_in, 256) + temb = time_embedding_mlp( + temb, inner_dim, name="time_text_embed.timestep_embedder" + ) + temb = temb + time_embedding_mlp( + pooled_in, inner_dim, name="time_text_embed.text_embedder" + ) + context = layers.Dense(caption_projection_dim, name="context_embedder")( + context_in + ) + for i in range(num_layers): + block = StableDiffusion3JointTransformerBlock( + inner_dim, + num_attention_heads, + module_path=f"transformer_blocks.{i}", + context_pre_only=i == num_layers - 1, + qk_norm=qk_norm, + use_dual_attention=i in dual_attention_layers, + ) + if i == num_layers - 1: + x = block([x, context, temb]) + else: + x, context = block([x, context, temb]) + x = StableDiffusion3AdaLayerNorm(inner_dim, 2, module_path="norm_out")( + [x, temb] + ) + x = layers.Dense(patch_size * patch_size * out_channels, name="proj_out")(x) + + # unpatchify: (B, h*w, p*p*C) -> (B, h*p, w*p, C) + h, w = sample_h // patch_size, sample_w // patch_size + x = ops.reshape(x, (-1, h, w, patch_size, patch_size, out_channels)) + x = ops.transpose(x, (0, 1, 3, 2, 4, 5)) + x = ops.reshape(x, (-1, h * patch_size, w * patch_size, out_channels)) + if data_format == "channels_first": + x = ops.transpose(x, (0, 3, 1, 2)) + + super().__init__( + inputs={ + "sample": sample_in, + "timestep": timestep_in, + "encoder_hidden_states": context_in, + "pooled_projections": pooled_in, + }, + outputs={"sample": x}, + name=name, + **kwargs, + ) + + self.data_format = data_format + self.channels_axis = channels_axis + self.sample_size = sample_size + self.patch_size = patch_size + self.in_channels = in_channels + self.out_channels = out_channels + self.num_layers = num_layers + self.attention_head_dim = attention_head_dim + self.num_attention_heads = num_attention_heads + self.joint_attention_dim = joint_attention_dim + self.caption_projection_dim = caption_projection_dim + self.pooled_projection_dim = pooled_projection_dim + self.pos_embed_max_size = pos_embed_max_size + self.qk_norm = qk_norm + self.dual_attention_layers = dual_attention_layers + self.text_seq_len = text_seq_len + + @classmethod + def from_diffusers_config(cls, config, **kwargs): + """Build from a diffusers ``transformer/config.json`` dict.""" + return cls(**cls.kwargs_from_diffusers_config(config), **kwargs) + + @staticmethod + def kwargs_from_diffusers_config(config): + """The constructor kwargs a diffusers ``transformer/config.json`` describes.""" + return { + "sample_size": config.get("sample_size", 128), + "patch_size": config.get("patch_size", 2), + "in_channels": config.get("in_channels", 16), + "out_channels": config.get("out_channels") or config.get("in_channels", 16), + "num_layers": config.get("num_layers", 24), + "attention_head_dim": config.get("attention_head_dim", 64), + "num_attention_heads": config.get("num_attention_heads", 24), + "joint_attention_dim": config.get("joint_attention_dim", 4096), + "caption_projection_dim": config.get("caption_projection_dim", 1536), + "pooled_projection_dim": config.get("pooled_projection_dim", 2048), + "pos_embed_max_size": config.get("pos_embed_max_size", 192), + "qk_norm": config.get("qk_norm"), + "dual_attention_layers": tuple(config.get("dual_attention_layers") or ()), + } + + def get_config(self): + config = super().get_config() + config.update( + { + "sample_size": self.sample_size, + "patch_size": self.patch_size, + "in_channels": self.in_channels, + "out_channels": self.out_channels, + "num_layers": self.num_layers, + "attention_head_dim": self.attention_head_dim, + "num_attention_heads": self.num_attention_heads, + "joint_attention_dim": self.joint_attention_dim, + "caption_projection_dim": self.caption_projection_dim, + "pooled_projection_dim": self.pooled_projection_dim, + "pos_embed_max_size": self.pos_embed_max_size, + "qk_norm": self.qk_norm, + "dual_attention_layers": self.dual_attention_layers, + "text_seq_len": self.text_seq_len, + "name": self.name, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class SD3T5EncoderModel(BaseModel): + """The third Stable Diffusion 3 text encoder: the encoder half of T5 v1.1 XXL + (transformers ``T5EncoderModel`` on ``google/t5-v1_1-xxl``), pure Keras 3. + + A shared token embedding, ``num_layers`` pre-norm blocks (bias-free self-attention + without ``1/sqrt(d)`` scaling, one learned relative position bias shared by all + blocks, the gated-GELU feed-forward of T5 v1.1) and a final RMSNorm. The + attention, norm and position-bias layers are the t5 family's; the gated + feed-forward, which the original T5 does not have, is + :class:`~zeromodels.models.stable_diffusion_3.stable_diffusion_3_layers.StableDiffusion3T5GatedFeedForward`. + Inputs ``{"input_ids", "attention_mask"}`` (``(B, T)`` int32), output + ``{"last_hidden_state": (B, T, embed_dim)}``. + + Every SD 3 / 3.5 checkpoint conditions on the same 4.7B-parameter encoder, so it + is hosted once (``zeromodels/t5-v1_1-xxl-encoder``, float16 weights) and attached + to a task model on demand: + ``StableDiffusion3TextToImage.from_weights(repo, text_encoder_3="zeromodels/t5-v1_1-xxl-encoder")``. + + Args: + vocab_size: SentencePiece vocabulary (32128). + embed_dim: Token width (4096). + key_value_dim: Width of one attention head (64). + mlp_dim: Feed-forward inner width (10240). + num_layers: Encoder blocks (24). + num_heads: Attention heads (64). + relative_attention_num_buckets: Relative-position-bias buckets (32). + relative_attention_max_distance: Largest bucketed distance (128). + layer_norm_eps: RMSNorm epsilon. + pad_token_id / eos_token_id: The tokenizer's ```` and ```` ids. + """ + + HF_MODEL_TYPE = None + config_class = StableDiffusion3T5EncoderConfig + + def __init__( + self, + vocab_size=32128, + embed_dim=4096, + key_value_dim=64, + mlp_dim=10240, + num_layers=24, + num_heads=64, + relative_attention_num_buckets=32, + relative_attention_max_distance=128, + layer_norm_eps=1e-6, + pad_token_id=0, + eos_token_id=1, + name="SD3T5EncoderModel", + **kwargs, + ): + shared = layers.Embedding(vocab_size, embed_dim, name="shared") + rel_bias = layers.Embedding( + relative_attention_num_buckets, num_heads, name="encoder_rel_bias" + ) + bias = T5PositionBias( + rel_bias, + num_heads, + relative_attention_num_buckets, + relative_attention_max_distance, + bidirectional=True, + causal=False, + name="encoder_bias", + ) + blocks = [ + StableDiffusion3T5GatedEncoderBlock( + embed_dim, + key_value_dim, + num_heads, + mlp_dim, + layer_norm_eps, + prefix=f"enc_{i}", + name=f"encoder_block_{i}", + ) + for i in range(num_layers) + ] + final_norm = T5LayerNorm(layer_norm_eps, name="encoder_final_layer_norm") + + input_ids_in = layers.Input(shape=(None,), dtype="int32", name="input_ids") + attn_in = layers.Input(shape=(None,), dtype="int32", name="attention_mask") + hidden = shared(input_ids_in) + position_bias = bias(hidden, attn_in) + for block in blocks: + hidden = block(hidden, position_bias) + hidden = final_norm(hidden) + super().__init__( + inputs={"input_ids": input_ids_in, "attention_mask": attn_in}, + outputs={"last_hidden_state": hidden}, + name=name, + **kwargs, + ) + self.shared = shared + self.encoder_rel_bias = rel_bias + self.encoder_bias = bias + self.encoder_blocks = blocks + self.encoder_final_layer_norm = final_norm + self.vocab_size = vocab_size + self.embed_dim = embed_dim + self.key_value_dim = key_value_dim + self.mlp_dim = mlp_dim + self.num_layers = num_layers + self.num_heads = num_heads + self.relative_attention_num_buckets = relative_attention_num_buckets + self.relative_attention_max_distance = relative_attention_max_distance + self.layer_norm_eps = layer_norm_eps + self.pad_token_id = pad_token_id + self.eos_token_id = eos_token_id + # the relative bias table is read inside the position-bias layer at run + # time (not traced into the graph): one eager call creates its weight so + # the checkpoint can load into it + with inference_scope(): + self( + { + "input_ids": ops.zeros((1, 4), dtype="int32"), + "attention_mask": ops.ones((1, 4), dtype="int32"), + } + ) + + @classmethod + def from_hf_config(cls, hf_config, **kwargs): + """Build from a transformers ``T5Config`` dict (``text_encoder_3/config.json``).""" + return cls(**cls.kwargs_from_hf_config(hf_config), **kwargs) + + @staticmethod + def kwargs_from_hf_config(hf_config): + """The constructor kwargs a transformers T5 v1.1 ``config.json`` describes.""" + return { + "vocab_size": hf_config.get("vocab_size", 32128), + "embed_dim": hf_config.get("d_model", 4096), + "key_value_dim": hf_config.get("d_kv", 64), + "mlp_dim": hf_config.get("d_ff", 10240), + "num_layers": hf_config.get("num_layers", 24), + "num_heads": hf_config.get("num_heads", 64), + "relative_attention_num_buckets": hf_config.get( + "relative_attention_num_buckets", 32 + ), + "relative_attention_max_distance": hf_config.get( + "relative_attention_max_distance", 128 + ), + "layer_norm_eps": hf_config.get("layer_norm_epsilon", 1e-6), + "pad_token_id": hf_config.get("pad_token_id", 0), + "eos_token_id": hf_config.get("eos_token_id", 1), + } + + def get_config(self): + config = super().get_config() + config.update( + { + "vocab_size": self.vocab_size, + "embed_dim": self.embed_dim, + "key_value_dim": self.key_value_dim, + "mlp_dim": self.mlp_dim, + "num_layers": self.num_layers, + "num_heads": self.num_heads, + "relative_attention_num_buckets": self.relative_attention_num_buckets, + "relative_attention_max_distance": self.relative_attention_max_distance, + "layer_norm_eps": self.layer_norm_eps, + "pad_token_id": self.pad_token_id, + "eos_token_id": self.eos_token_id, + "name": self.name, + } + ) + return config + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3Model(StableDiffusionModel): + """Stable Diffusion 3 weights container: the MMDiT, the 16-channel VAE and the + two CLIP text encoders (ViT-L/14 and OpenCLIP ViT-bigG/14, each with its + projection) as one functional ``keras.Model``. + + The SD 1.x container (:class:`StableDiffusionModel`) with the SD 3 parts: four + disconnected sub-graphs, ``{"sample", "timestep", "encoder_hidden_states", + "pooled_projections"} -> "noise_pred"`` (the transformer), ``{"image", + "latent"} -> "moments" / "image"`` (the VAE) and ``{"token_ids", "token_ids_2", + "padding_mask"} -> "prompt_embeds" / "pooled_prompt_embeds"`` (the CLIP towers: + penultimate states concatenated to 2048-d, projected pooled states concatenated + to 2048-d). Components: ``.transformer`` / ``.vae`` / ``.text_encoder`` / + ``.text_encoder_2``. + + The third text encoder, the 4.7B-parameter T5-XXL, is not part of the + container: it is shared by every SD 3 / 3.5 checkpoint and hosted once + (``zeromodels/t5-v1_1-xxl-encoder``, an :class:`SD3T5EncoderModel`), and the + task class attaches it on demand (``text_encoder_3``); without it the + T5 features are zeros, SD 3's documented memory-saving mode. Hosted as one + ``zm_config.json`` + float16 weights per variant; on-the-fly ``hf:`` conversion + is not supported for diffusion models. + + Args: + **kwargs: The flat :class:`StableDiffusion3Config` fields + (``transformer_`` / ``vae_`` / ``text_`` / ``text_2_`` prefixed), or a + config positionally. + """ + + config_class = StableDiffusion3Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_3_HUB_SIBLINGS + + def __init__(self, name="StableDiffusion3Model", **kwargs): + super().__init__(name=name, **kwargs) + + def build_components(self, config, data_format, channels_axis): + d, v, t, t2 = ( + config.transformer_config, + config.vae_config, + config.text_config, + config.text_config_2, + ) + transformer = SD3Transformer2DModel( + d, + text_seq_len=t.max_seq_len + config.max_sequence_length, + data_format=data_format, + channels_axis=channels_axis, + ) + vae = AutoencoderKL( + **v.constructor_kwargs(), + data_format=data_format, + channels_axis=channels_axis, + ) + text_encoder = build_text_encoder( + t, + config.hidden_act, + config.layer_norm_eps, + projection_dim=t.projection_dim, + name="text_encoder", + ) + text_encoder_2 = build_text_encoder( + t2, + t2.hidden_act, + config.layer_norm_eps, + projection_dim=t2.projection_dim, + name="text_encoder_2", + ) + return { + "transformer": transformer, + "vae": vae, + "text_encoder": text_encoder, + "text_encoder_2": text_encoder_2, + } + + def build_graph(self, config, components, data_format): + d, v, t = config.transformer_config, config.vae_config, config.text_config + transformer, vae = components["transformer"], components["vae"] + text_encoder, text_encoder_2 = ( + components["text_encoder"], + components["text_encoder_2"], + ) + sample_h, sample_w = ( + d.sample_size + if isinstance(d.sample_size, (tuple, list)) + else (d.sample_size, d.sample_size) + ) + img_h, img_w = ( + v.sample_size + if isinstance(v.sample_size, (tuple, list)) + else (v.sample_size, v.sample_size) + ) + lat_h, lat_w = img_h // vae.vae_scale_factor, img_w // vae.vae_scale_factor + inputs = { + "sample": layers.Input( + shape=(d.in_channels, sample_h, sample_w) + if data_format == "channels_first" + else (sample_h, sample_w, d.in_channels), + name="sample", + ), + "timestep": layers.Input(shape=(), name="timestep"), + "encoder_hidden_states": layers.Input( + shape=(transformer.text_seq_len, d.joint_attention_dim), + name="encoder_hidden_states", + ), + "pooled_projections": layers.Input( + shape=(d.pooled_projection_dim,), name="pooled_projections" + ), + "image": layers.Input( + shape=(v.in_channels, img_h, img_w) + if data_format == "channels_first" + else (img_h, img_w, v.in_channels), + name="image", + ), + "latent": layers.Input( + shape=(v.latent_channels, lat_h, lat_w) + if data_format == "channels_first" + else (lat_h, lat_w, v.latent_channels), + name="latent", + ), + "token_ids": layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="token_ids" + ), + "token_ids_2": layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="token_ids_2" + ), + "padding_mask": layers.Input( + shape=(t.max_seq_len,), dtype="int32", name="padding_mask" + ), + } + noise_pred = transformer( + { + key: inputs[key] + for key in ( + "sample", + "timestep", + "encoder_hidden_states", + "pooled_projections", + ) + } + )["sample"] + vae_out = vae({"image": inputs["image"], "latent": inputs["latent"]}) + text_out = text_encoder( + {"token_ids": inputs["token_ids"], "padding_mask": inputs["padding_mask"]} + ) + text_out_2 = text_encoder_2( + { + "token_ids": inputs["token_ids_2"], + "padding_mask": inputs["padding_mask"], + } + ) + outputs = { + "noise_pred": noise_pred, + "moments": vae_out["moments"], + "image": vae_out["sample"], + "prompt_embeds": ops.concatenate( + [ + text_out["penultimate_hidden_state"], + text_out_2["penultimate_hidden_state"], + ], + axis=-1, + ), + "pooled_prompt_embeds": ops.concatenate( + [text_out["text_embeds"], text_out_2["text_embeds"]], axis=-1 + ), + } + return inputs, outputs + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3TextToImage(StableDiffusionTextToImage, StableDiffusion3Model): + """Text-to-image Stable Diffusion 3, pure Keras 3 and cross-backend (the + diffusers ``StableDiffusion3Pipeline`` in this library's task-class form). + + :class:`StableDiffusionTextToImage`'s ``generate`` (``BaseDiffusion``) over the + :class:`StableDiffusion3Model` graph with the rectified-flow sampler + (:class:`FlowMatchEulerDiscreteScheduler`). The prompt goes through the two + CLIP towers (penultimate states concatenated, zero-padded to the T5 width) and, + when a T5-XXL encoder is attached, through it (256 tokens); the text tokens are + concatenated along the sequence, the pooled CLIP states along the features. With + no negative prompt the empty prompt is encoded (no zeroing). The T5 encoder is + loaded separately and optionally:: + + sd = StableDiffusion3TextToImage.from_weights( + "zeromodels/stable-diffusion-3-medium", + text_encoder_3="zeromodels/t5-v1_1-xxl-encoder", # or omit: T5 features zeroed + ) + tokenizer = StableDiffusion3Tokenizer.from_weights("zeromodels/stable-diffusion-3-medium") + image = sd.generate(**tokenizer("a photo of a cat")) # (1, 1024, 1024, 3) uint8 + + ``sd.text_encoder_3`` can also be assigned directly (an + :class:`SD3T5EncoderModel`, e.g. one loaded with ``quantization="int8"``); it + lives outside the container's weights. + + Args: + scheduler: A :class:`BaseScheduler`; defaults to the config's + ``scheduler_config``, else the flow-match Euler sampler with shift 3. + **kwargs: The flat :class:`StableDiffusion3Config` fields. + """ + + config_class = StableDiffusion3Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_3_HUB_SIBLINGS + # diffusers' StableDiffusion3Pipeline defaults; a repo's generate_args override + generate_args = {"num_inference_steps": 28, "guidance_scale": 7.0} + + def __init__(self, scheduler=None, name="StableDiffusion3TextToImage", **kwargs): + super().__init__(scheduler=scheduler, name=name, **kwargs) + + @classmethod + def from_weights(cls, identifier, text_encoder_3=None, **kwargs): + """``BaseModel.from_weights`` plus ``text_encoder_3``: a hosted + :class:`SD3T5EncoderModel` repo id (loaded with the same ``load_dtype``) or + a built model to attach as the third text encoder.""" + model = super().from_weights(identifier, **kwargs) + if text_encoder_3 is not None: + if isinstance(text_encoder_3, str): + text_encoder_3 = SD3T5EncoderModel.from_weights( + text_encoder_3, load_dtype=kwargs.get("load_dtype") + ) + model.text_encoder_3 = text_encoder_3 + return model + + @property + def text_encoder_3(self): + return self.__dict__.get("_text_encoder_3") + + def __setattr__(self, name, value): + # the T5 encoder is kept out of the tracked (saved / loaded / device-moved) + # sub-layers, like the LM mixins' caches: Keras (and torch's Module on that + # backend) would otherwise register any model assigned as an attribute + if name == "text_encoder_3": + self.__dict__["_text_encoder_3"] = value + return + super().__setattr__(name, value) + + def default_scheduler(self): + return FlowMatchEulerDiscreteScheduler(shift=3.0) + + @property + def latent_shape(self): + size = self.transformer.sample_size + height, width = size if isinstance(size, (tuple, list)) else (size, size) + if self.data_format == "channels_first": + return (self.transformer.in_channels, height, width) + return (height, width, self.transformer.in_channels) + + def unconditional_ids(self, batch): + cfg = self.config + length = cfg.text_config.max_seq_len + row = [cfg.bos_token_id, cfg.eos_token_id] + [cfg.pad_token_id] * (length - 2) + return ops.convert_to_tensor([row] * batch, dtype="int32") + + def unconditional_ids_3(self, batch): + cfg = self.config + row = [cfg.eos_token_id_3] + [cfg.pad_token_id_3] * ( + cfg.max_sequence_length - 1 + ) + return ops.convert_to_tensor([row] * batch, dtype="int32") + + def prompt_mask(self, input_ids): + # the tokens up to and including the first <|endoftext|>; what follows is + # padding (the first tokenizer pads with <|endoftext|> itself) + is_eos = ops.cast(ops.equal(input_ids, self.config.eos_token_id), "int32") + return ops.equal(ops.cumsum(is_eos, axis=1) - is_eos, 0) + + def encode_prompt(self, input_ids, attention_mask=None, input_ids_3=None): + input_ids = ops.cast(ops.convert_to_tensor(input_ids), "int32") + batch = int(input_ids.shape[0]) + if attention_mask is None: + real = self.prompt_mask(input_ids) + else: + real = ops.greater( + ops.cast(ops.convert_to_tensor(attention_mask), "int32"), 0 + ) + # the second tokenizer pads with "!" (id 0) where the first pads with + # <|endoftext|>; the towers attend the padding, as the reference does + input_ids_2 = ops.where(real, input_ids, self.config.pad_token_id_2) + ones = ops.ones_like(input_ids) + out = self.text_encoder({"token_ids": input_ids, "padding_mask": ones}) + out_2 = self.text_encoder_2({"token_ids": input_ids_2, "padding_mask": ones}) + clip_embeds = ops.concatenate( + [out["penultimate_hidden_state"], out_2["penultimate_hidden_state"]], + axis=-1, + ) + pooled = ops.concatenate([out["text_embeds"], out_2["text_embeds"]], axis=-1) + # the CLIP features are zero-padded to the T5 width, then the T5 tokens + # (or zeros without a T5 encoder) follow along the sequence + t5_dim = self.transformer.joint_attention_dim + clip_embeds = ops.pad( + clip_embeds, ((0, 0), (0, 0), (0, t5_dim - clip_embeds.shape[-1])) + ) + length = self.config.max_sequence_length + if self.text_encoder_3 is None: + t5_embeds = ops.zeros((batch, length, t5_dim), dtype=clip_embeds.dtype) + else: + if input_ids_3 is None: + input_ids_3 = self.unconditional_ids_3(batch) + input_ids_3 = ops.cast(ops.convert_to_tensor(input_ids_3), "int32") + if int(input_ids_3.shape[0]) == 1 and batch > 1: + input_ids_3 = ops.repeat(input_ids_3, batch, axis=0) # one for all + t5_embeds = self.text_encoder_3( + {"input_ids": input_ids_3, "attention_mask": ops.ones_like(input_ids_3)} + )["last_hidden_state"] + t5_embeds = ops.cast(t5_embeds, clip_embeds.dtype) + return { + "encoder_hidden_states": ops.concatenate([clip_embeds, t5_embeds], axis=1), + "pooled_projections": pooled, + } + + def encode_negative_prompt(self, negative_input_ids, batch, **conditioning): + if negative_input_ids is None: + # the empty prompt for the T5 tower as well, not the positive prompt's ids + conditioning = dict(conditioning, input_ids_3=None) + return super().encode_negative_prompt(negative_input_ids, batch, **conditioning) + + def predict_noise(self, latents, timesteps, embeddings): + return self.transformer( + {"sample": latents, "timestep": timesteps, **embeddings} + )["sample"] diff --git a/zeromodels/models/stable_diffusion_3/stable_diffusion_3_tokenizer.py b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_tokenizer.py new file mode 100644 index 00000000..ee4ebb65 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3/stable_diffusion_3_tokenizer.py @@ -0,0 +1,89 @@ +import keras +from keras import ops + +from zeromodels.models.stable_diffusion.stable_diffusion_tokenizer import ( + StableDiffusionTokenizer, +) +from zeromodels.models.t5.t5_tokenizer import T5Tokenizer + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3Tokenizer(StableDiffusionTokenizer): + """Stable Diffusion 3 text tokenizer: the CLIP ViT-L/14 BPE tokenizer plus the + T5 SentencePiece tokenizer, in one call. + + SD 3 tokenizes the prompt three times: with the two CLIP tokenizers (the same + BPE, padded with ``<|endoftext|>`` and ``!`` respectively; this class is the + first, and :class:`StableDiffusion3TextToImage` derives the second tower's ids + from the returned ``attention_mask``, as SDXL does) and with the T5 tokenizer + (````-terminated, ````-padded to ``max_sequence_length``, 256). The + returned dict is ``{"input_ids", "attention_mask", "input_ids_3"}``, what + ``generate`` takes. Loads by repo id like the model, reading the hosted repo's + ``tokenizer.json`` (CLIP) and ``tokenizer_3.json`` (T5): + ``from_weights("zeromodels/stable-diffusion-3-medium")``. + + Args: + hf_id: Hosted repo to read the tokenizer files from. Required unless the + files are given (no default repo). + tokenizer_file: Explicit CLIP ``tokenizer.json`` path (overrides ``hf_id``). + tokenizer_file_3: Explicit T5 ``tokenizer.json`` path (overrides ``hf_id``). + max_seq_len: CLIP padded / truncated length (default 77). + max_sequence_length: T5 padded / truncated length (default 256). + pad_token: CLIP pad token string (``<|endoftext|>``). + """ + + def __init__( + self, + hf_id=None, + tokenizer_file=None, + tokenizer_file_3=None, + max_seq_len=77, + max_sequence_length=256, + pad_token="<|endoftext|>", + **kwargs, + ): + super().__init__( + hf_id=hf_id, + tokenizer_file=tokenizer_file, + max_seq_len=max_seq_len, + pad_token=pad_token, + **kwargs, + ) + if tokenizer_file_3 is None: + tokenizer_file_3 = self.download_tokenizer_json(hf_id, "tokenizer_3.json") + self.tokenizer_file_3 = tokenizer_file_3 + self.max_sequence_length = max_sequence_length + self.tokenizer_3 = T5Tokenizer( + tokenizer_file=tokenizer_file_3, max_seq_len=max_sequence_length + ) + + @staticmethod + def download_tokenizer_json(hf_id, filename="tokenizer.json"): + import os + + from huggingface_hub import hf_hub_download + + return hf_hub_download(hf_id, filename, token=os.environ.get("HF_TOKEN")) + + def call(self, inputs): + out = super().call(inputs) + # the T5 tokenizer truncates to max_sequence_length and pads to the longest + # prompt; the transformer takes the fixed 256-token length, so pad the rest + ids = self.tokenizer_3(inputs)["input_ids"] + short = self.max_sequence_length - int(ids.shape[1]) + if short > 0: + ids = ops.pad( + ids, ((0, 0), (0, short)), constant_values=self.tokenizer_3.pad_token_id + ) + out["input_ids_3"] = ids + return out + + def get_config(self): + config = super().get_config() + config.update( + { + "tokenizer_file_3": self.tokenizer_file_3, + "max_sequence_length": self.max_sequence_length, + } + ) + return config diff --git a/zeromodels/models/stable_diffusion_3_5/__init__.py b/zeromodels/models/stable_diffusion_3_5/__init__.py new file mode 100644 index 00000000..29ea0ff8 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3_5/__init__.py @@ -0,0 +1,17 @@ +from .stable_diffusion_3_5_config import ( + StableDiffusion3_5Config, + StableDiffusion3_5TransformerConfig, +) +from .stable_diffusion_3_5_model import ( + StableDiffusion3_5Model, + StableDiffusion3_5TextToImage, +) +from .stable_diffusion_3_5_tokenizer import StableDiffusion3_5Tokenizer + +__all__ = [ + "StableDiffusion3_5Model", + "StableDiffusion3_5TextToImage", + "StableDiffusion3_5Config", + "StableDiffusion3_5TransformerConfig", + "StableDiffusion3_5Tokenizer", +] diff --git a/zeromodels/models/stable_diffusion_3_5/convert_stable_diffusion_3_5_diffusers_to_keras.py b/zeromodels/models/stable_diffusion_3_5/convert_stable_diffusion_3_5_diffusers_to_keras.py new file mode 100644 index 00000000..52c919c3 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3_5/convert_stable_diffusion_3_5_diffusers_to_keras.py @@ -0,0 +1,79 @@ +import numpy as np + +from zeromodels.models.stable_diffusion_3.convert_stable_diffusion_3_diffusers_to_keras import ( + config_from_diffusers as stable_diffusion_3_config_from_diffusers, +) +from zeromodels.models.stable_diffusion_3.convert_stable_diffusion_3_diffusers_to_keras import ( + transfer_stable_diffusion_3, +) + +STABLE_DIFFUSION_3_5_SOURCES = { + "stable-diffusion-3.5-large": "stabilityai/stable-diffusion-3.5-large", + "stable-diffusion-3.5-large-turbo": "stabilityai/stable-diffusion-3.5-large-turbo", + "stable-diffusion-3.5-medium": "stabilityai/stable-diffusion-3.5-medium", +} + + +def config_from_diffusers(repo, token=None): + from zeromodels.models.stable_diffusion_3_5.stable_diffusion_3_5_config import ( + StableDiffusion3_5Config, + ) + + return stable_diffusion_3_config_from_diffusers( + repo, token=token, config_cls=StableDiffusion3_5Config + ) + + +def transfer_stable_diffusion_3_5(repo, token=None, dtype="float16"): + from zeromodels.models.stable_diffusion_3_5.stable_diffusion_3_5_config import ( + StableDiffusion3_5Config, + ) + from zeromodels.models.stable_diffusion_3_5.stable_diffusion_3_5_model import ( + StableDiffusion3_5Model, + ) + + return transfer_stable_diffusion_3( + repo, + token=token, + dtype=dtype, + model_cls=StableDiffusion3_5Model, + config_cls=StableDiffusion3_5Config, + ) + + +if __name__ == "__main__": + import gc + import os + + import keras + + OUT_DIR = os.environ.get("ZM_OUT_DIR", "C:/Users/gites/Desktop/code/v1_weights") + os.makedirs(OUT_DIR, exist_ok=True) + MAX_SHARD_GB = 5.0 + token = os.environ.get("HF_TOKEN") + selected = [v for v in os.environ.get("ZM_VARIANTS", "").split(",") if v] + sources = { + variant: source + for variant, source in STABLE_DIFFUSION_3_5_SOURCES.items() + if not selected or variant in selected + } + + for variant, source in sources.items(): + print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") + model, config = transfer_stable_diffusion_3_5(source, token=token) + + n_bytes = sum( + int(np.prod(w.shape)) * np.dtype(w.dtype).itemsize for w in model.weights + ) + stem = os.path.join(OUT_DIR, variant.replace("-", "_").replace(".", "_")) + if n_bytes > MAX_SHARD_GB * 1024**3: + out = f"{stem}.weights.json" + model.save_weights(out, max_shard_size=MAX_SHARD_GB) + else: + out = f"{stem}.weights.h5" + model.save_weights(out) + print(f" saved -> {out} ({n_bytes / 1024**3:.2f} GB)") + + del model + keras.backend.clear_session() + gc.collect() diff --git a/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_config.py b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_config.py new file mode 100644 index 00000000..f71459a8 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_config.py @@ -0,0 +1,52 @@ +from zeromodels.models.stable_diffusion_3.stable_diffusion_3_config import ( + StableDiffusion3Config, + StableDiffusion3TransformerConfig, +) + + +class StableDiffusion3_5TransformerConfig(StableDiffusion3TransformerConfig): + r"""The Stable Diffusion 3.5 MMDiT: a [`StableDiffusion3TransformerConfig`] with + the SD 3.5 large layout. + + Same fields as [`StableDiffusion3TransformerConfig`]; the defaults are SD 3.5 + large (8B parameters): 38 blocks of 38 heads (2432 wide) with RMS-normalized + queries and keys. SD 3.5 medium (2.5B, MMDiT-X) is 24 blocks of 24 heads with + ``dual_attention_layers`` 0 to 12 and a 384-patch position grid; each repo's + ``zm_config.json`` carries its own values.""" + + num_layers: int = 38 + num_attention_heads: int = 38 + caption_projection_dim: int = 2432 + qk_norm: str | None = "rms_norm" + + +class StableDiffusion3_5Config(StableDiffusion3Config): + r"""Configuration for [`StableDiffusion3_5Model`], the hosted SD 3.5 weights container. + + Same shape as [`StableDiffusion3Config`] (nested ``transformer_config`` / + ``vae_config`` / ``text_config`` / ``text_config_2`` plus the scheduler, the + token ids and the T5 length; flat constructor) with the + [`StableDiffusion3_5TransformerConfig`] denoiser. The three SD 3.5 checkpoints + (large, large-turbo, medium) share the VAE and the text encoders with SD 3. + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_3_5 import StableDiffusion3_5Config, StableDiffusion3_5Model + + >>> config = StableDiffusion3_5Config() + >>> model = StableDiffusion3_5Model(config) # random weights, SD 3.5 large shapes + >>> model.config.transformer_config.qk_norm + 'rms_norm' + ```""" + + model_type = "stable_diffusion_3_5" + + sub_configs = { + "transformer_config": StableDiffusion3_5TransformerConfig, + "vae_config": StableDiffusion3Config.sub_configs["vae_config"], + "text_config": StableDiffusion3Config.sub_configs["text_config"], + "text_config_2": StableDiffusion3Config.sub_configs["text_config_2"], + } + + transformer_config: StableDiffusion3_5TransformerConfig | dict | None = None diff --git a/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_model.py b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_model.py new file mode 100644 index 00000000..12c56582 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_model.py @@ -0,0 +1,69 @@ +import keras + +from zeromodels.models.stable_diffusion_3.stable_diffusion_3_model import ( + StableDiffusion3Model, + StableDiffusion3TextToImage, +) + +from .stable_diffusion_3_5_config import StableDiffusion3_5Config + +# The container and the task build the identical graph, so both load the one +# hosted repo (whose zm_config.json names the container). +STABLE_DIFFUSION_3_5_HUB_SIBLINGS = frozenset( + {"StableDiffusion3_5Model", "StableDiffusion3_5TextToImage"} +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3_5Model(StableDiffusion3Model): + """Stable Diffusion 3.5 weights container: the MMDiT (large) or MMDiT-X (medium), + the 16-channel VAE and the two CLIP text encoders as one functional + ``keras.Model``. + + The SD 3 container (:class:`StableDiffusion3Model`) with the SD 3.5 + configuration (:class:`StableDiffusion3_5Config`): RMS-normalized queries and + keys in every block and, for the medium checkpoint, the dual-attention blocks. + Same components (``.transformer`` / ``.vae`` / ``.text_encoder`` / + ``.text_encoder_2``), same separately hosted T5-XXL. Hosted under + ``zeromodels/stable-diffusion-3.5-*`` in float16. + + Args: + **kwargs: The flat :class:`StableDiffusion3_5Config` fields, or a config + positionally. + """ + + config_class = StableDiffusion3_5Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_3_5_HUB_SIBLINGS + + def __init__(self, name="StableDiffusion3_5Model", **kwargs): + super().__init__(name=name, **kwargs) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3_5TextToImage(StableDiffusion3TextToImage): + """Text-to-image Stable Diffusion 3.5, pure Keras 3 and cross-backend. + + :class:`StableDiffusion3TextToImage` with the SD 3.5 configuration: the same + ``generate`` over the SD 3.5 graph, the repo's flow-match sampler and + defaults (28 steps at guidance 3.5 for large, 40 at 4.5 for medium, 4 without + guidance for large-turbo):: + + sd = StableDiffusion3_5TextToImage.from_weights( + "zeromodels/stable-diffusion-3.5-large", text_encoder_3="zeromodels/t5-v1_1-xxl-encoder" + ) + tokenizer = StableDiffusion3_5Tokenizer.from_weights("zeromodels/stable-diffusion-3.5-large") + image = sd.generate(**tokenizer("a photo of a cat")) # (1, 1024, 1024, 3) uint8 + + Args: + scheduler: A :class:`BaseScheduler`; defaults to the config's + ``scheduler_config``. + **kwargs: The flat :class:`StableDiffusion3_5Config` fields. + """ + + config_class = StableDiffusion3_5Config + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_3_5_HUB_SIBLINGS + # diffusers' defaults for SD 3.5 large; each repo's generate_args override + generate_args = {"num_inference_steps": 28, "guidance_scale": 3.5} + + def __init__(self, scheduler=None, name="StableDiffusion3_5TextToImage", **kwargs): + super().__init__(scheduler=scheduler, name=name, **kwargs) diff --git a/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_tokenizer.py b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_tokenizer.py new file mode 100644 index 00000000..b3259457 --- /dev/null +++ b/zeromodels/models/stable_diffusion_3_5/stable_diffusion_3_5_tokenizer.py @@ -0,0 +1,39 @@ +import keras + +from zeromodels.models.stable_diffusion_3.stable_diffusion_3_tokenizer import ( + StableDiffusion3Tokenizer, +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusion3_5Tokenizer(StableDiffusion3Tokenizer): + """Stable Diffusion 3.5 text tokenizer: the SD 3 tokenizer (CLIP BPE plus the T5 + SentencePiece tokenizer, ``{"input_ids", "attention_mask", "input_ids_3"}``); + the checkpoints share their tokenizers with SD 3. Loads by repo id like the + model: ``from_weights("zeromodels/stable-diffusion-3.5-large")``. + + Args: + hf_id: Hosted repo to read the tokenizer files from. Required unless the + files are given (no default repo). + tokenizer_file / tokenizer_file_3: Explicit CLIP / T5 ``tokenizer.json`` paths. + max_seq_len: CLIP padded / truncated length (default 77). + max_sequence_length: T5 padded / truncated length (default 256). + """ + + def __init__( + self, + hf_id=None, + tokenizer_file=None, + tokenizer_file_3=None, + max_seq_len=77, + max_sequence_length=256, + **kwargs, + ): + super().__init__( + hf_id=hf_id, + tokenizer_file=tokenizer_file, + tokenizer_file_3=tokenizer_file_3, + max_seq_len=max_seq_len, + max_sequence_length=max_sequence_length, + **kwargs, + ) diff --git a/zeromodels/models/stable_diffusion_xl/__init__.py b/zeromodels/models/stable_diffusion_xl/__init__.py new file mode 100644 index 00000000..0943c298 --- /dev/null +++ b/zeromodels/models/stable_diffusion_xl/__init__.py @@ -0,0 +1,29 @@ +from .stable_diffusion_xl_config import ( + StableDiffusionXLConfig, + StableDiffusionXLRefinerConfig, + StableDiffusionXLRefinerUNetConfig, + StableDiffusionXLTextConfig2, + StableDiffusionXLUNetConfig, + StableDiffusionXLVAEConfig, +) +from .stable_diffusion_xl_model import ( + StableDiffusionXLModel, + StableDiffusionXLRefinerImageToImage, + StableDiffusionXLRefinerModel, + StableDiffusionXLTextToImage, +) +from .stable_diffusion_xl_tokenizer import StableDiffusionXLTokenizer + +__all__ = [ + "StableDiffusionXLModel", + "StableDiffusionXLTextToImage", + "StableDiffusionXLRefinerModel", + "StableDiffusionXLRefinerImageToImage", + "StableDiffusionXLConfig", + "StableDiffusionXLRefinerConfig", + "StableDiffusionXLRefinerUNetConfig", + "StableDiffusionXLTextConfig2", + "StableDiffusionXLUNetConfig", + "StableDiffusionXLVAEConfig", + "StableDiffusionXLTokenizer", +] diff --git a/zeromodels/models/stable_diffusion_xl/convert_stable_diffusion_xl_diffusers_to_keras.py b/zeromodels/models/stable_diffusion_xl/convert_stable_diffusion_xl_diffusers_to_keras.py new file mode 100644 index 00000000..b29628d7 --- /dev/null +++ b/zeromodels/models/stable_diffusion_xl/convert_stable_diffusion_xl_diffusers_to_keras.py @@ -0,0 +1,299 @@ +import collections.abc + +import numpy as np +from tqdm import tqdm + +from zeromodels.conversion.exceptions import ( + WeightMappingError, + WeightShapeMismatchError, +) +from zeromodels.conversion.weight_split_util import split_model_weights +from zeromodels.conversion.weight_transfer_util import ( + compare_keras_torch_names, + transfer_weights, +) +from zeromodels.models.clip.convert_clip_hf_to_keras import transfer_clip_weights +from zeromodels.models.stable_diffusion.convert_stable_diffusion_diffusers_to_keras import ( + WEIGHT_NAME_MAPPING, +) + +# SDXL ships in float16 (the release checkpoints' precision; the diffusers fp32 +# files are upcasts of the fp16 variant), so the container is built and saved in +# float16 (the VAE excepted: force_upcast keeps it float32). Loading a repo +# rebuilds at that precision by default; load_dtype="float32" upcasts. The 0.9 +# repos are gated (research license): set HF_TOKEN to convert them. The refiner +# repos (an image-to-image model: one text tower, a four-level UNet) convert to +# StableDiffusionXLRefinerModel, everything else to StableDiffusionXLModel. +STABLE_DIFFUSION_XL_SOURCES = { + "stable-diffusion-xl-base-0.9": "stabilityai/stable-diffusion-xl-base-0.9", + "stable-diffusion-xl-refiner-0.9": "stabilityai/stable-diffusion-xl-refiner-0.9", + "stable-diffusion-xl-base-1.0": "stabilityai/stable-diffusion-xl-base-1.0", + "stable-diffusion-xl-refiner-1.0": "stabilityai/stable-diffusion-xl-refiner-1.0", + "sdxl-turbo": "stabilityai/sdxl-turbo", +} + + +class LazyNumpyStateDict(collections.abc.Mapping): + def __init__(self, module): + self.state = module.state_dict() + + def __getitem__(self, key): + return self.state[key].detach().cpu().numpy() + + def __contains__(self, key): + return key in self.state + + def __iter__(self): + return iter(self.state) + + def __len__(self): + return len(self.state) + + +def text_config_kwargs(text): + hidden = text["hidden_size"] + return { + "hidden_dim": hidden, + "num_heads": text.get("num_attention_heads", 12), + "num_layers": text.get("num_hidden_layers", 12), + "mlp_ratio": text.get("intermediate_size", hidden * 4) / hidden, + "vocab_size": text.get("vocab_size", 49408), + "max_seq_len": text.get("max_position_embeddings", 77), + } + + +def config_from_diffusers(repo, token=None): + from diffusers import AutoencoderKL, EulerDiscreteScheduler, UNet2DConditionModel + from diffusers.pipelines.pipeline_utils import DiffusionPipeline + from transformers import CLIPTextConfig, CLIPTokenizerFast + + from zeromodels.models.stable_diffusion.stable_diffusion_model import ( + UNet2DConditionModel as KerasUNet2DConditionModel, + ) + from zeromodels.models.stable_diffusion_xl.stable_diffusion_xl_config import ( + StableDiffusionXLConfig, + StableDiffusionXLRefinerConfig, + ) + + pipeline = dict(DiffusionPipeline.load_config(repo, token=token)) + has_text = pipeline.get("text_encoder", [None])[0] is not None # refiner: no + force_zeros_for_empty_prompt = bool( + pipeline.get("force_zeros_for_empty_prompt", True) + ) + requires_aesthetics_score = bool(pipeline.get("requires_aesthetics_score", False)) + unet = dict(UNet2DConditionModel.load_config(repo, subfolder="unet", token=token)) + vae = dict(AutoencoderKL.load_config(repo, subfolder="vae", token=token)) + text = ( + CLIPTextConfig.from_pretrained( + repo, subfolder="text_encoder", token=token + ).to_dict() + if has_text + else None + ) + text_2 = CLIPTextConfig.from_pretrained( + repo, subfolder="text_encoder_2", token=token + ).to_dict() + scheduler = { + k: v + for k, v in EulerDiscreteScheduler.load_config( + repo, subfolder="scheduler", token=token + ).items() + if k == "_class_name" or not k.startswith("_") + } + tokenizer_2 = CLIPTokenizerFast.from_pretrained( + repo, subfolder="tokenizer_2", token=token + ) + tokenizer = ( + CLIPTokenizerFast.from_pretrained(repo, subfolder="tokenizer", token=token) + if has_text + else tokenizer_2 + ) + unet_kwargs = KerasUNet2DConditionModel.kwargs_from_diffusers_config(unet) + if requires_aesthetics_score: + # size (2) + crop (2) + aesthetic score (1) next to the pooled embedding + unet_kwargs["num_time_ids"] = 5 + config_cls = ( + StableDiffusionXLRefinerConfig if text is None else StableDiffusionXLConfig + ) + return config_cls( + unet_config=unet_kwargs, + vae_config={ + "in_channels": vae.get("in_channels", 3), + "out_channels": vae.get("out_channels", 3), + "latent_channels": vae.get("latent_channels", 4), + "block_out_channels": tuple(vae["block_out_channels"]), + "layers_per_block": vae.get("layers_per_block", 2), + "norm_num_groups": vae.get("norm_num_groups", 32), + "sample_size": unet.get("sample_size", 128) + * 2 ** (len(vae["block_out_channels"]) - 1), + "scaling_factor": vae.get("scaling_factor") or 0.13025, + "force_upcast": bool(vae.get("force_upcast", True)), + }, + text_config=None if text is None else text_config_kwargs(text), + text_config_2={ + **text_config_kwargs(text_2), + "projection_dim": text_2.get("projection_dim", 1280), + "hidden_act": text_2.get("hidden_act", "gelu"), + }, + hidden_act=(text or {}).get("hidden_act", "quick_gelu"), + layer_norm_eps=(text or text_2).get("layer_norm_eps", 1e-5), + scheduler_config=scheduler, + force_zeros_for_empty_prompt=force_zeros_for_empty_prompt, + requires_aesthetics_score=requires_aesthetics_score, + bos_token_id=tokenizer.bos_token_id, + eos_token_id=tokenizer.eos_token_id, + pad_token_id=tokenizer.pad_token_id, + pad_token_id_2=tokenizer_2.pad_token_id, + ) + + +def transfer_text_encoder(keras_model, hf_module): + state = { + k if k.startswith(("text_model.", "text_projection.")) else f"text_model.{k}": v + for k, v in { + k: v.detach().cpu().numpy() for k, v in hf_module.state_dict().items() + }.items() + } + transfer_clip_weights(keras_model, state) + + +def transfer_stable_diffusion_xl(repo, token=None, dtype="float16"): + import gc + + import torch + from diffusers import AutoencoderKL as DiffusersAutoencoderKL + from diffusers import UNet2DConditionModel as DiffusersUNet2DConditionModel + from transformers import CLIPTextModel, CLIPTextModelWithProjection + + from zeromodels.base.base_mixin import build_dtype_scope + from zeromodels.models.stable_diffusion_xl.stable_diffusion_xl_model import ( + StableDiffusionXLModel, + StableDiffusionXLRefinerModel, + ) + + config = config_from_diffusers(repo, token=token) + model_cls = ( + StableDiffusionXLRefinerModel + if config.text_config is None + else StableDiffusionXLModel + ) + with build_dtype_scope(dtype): + model = model_cls(config) + fp16 = {"torch_dtype": torch.float16, "variant": "fp16", "token": token} + fp32 = {"torch_dtype": torch.float32, "token": token} + + # the UNet (2.6B parameters: read tensor by tensor, no second in-memory copy) + # and the VAE: every keras weight's / maps to the diffusers key + # through WEIGHT_NAME_MAPPING. The SD 1.x era VAE files spell the mid-block + # attention query / key / value / proj_attn (diffusers renames them on load, a + # raw checkpoint read keeps them) + legacy = { + ".query.": ".to_q.", + ".key.": ".to_k.", + ".value.": ".to_v.", + ".proj_attn.": ".to_out.0.", + } + for component, module_cls, subfolder, load, lazy in ( + (model.unet, DiffusersUNet2DConditionModel, "unet", fp16, True), + (model.vae, DiffusersAutoencoderKL, "vae", fp32, False), + ): + module = module_cls.from_pretrained(repo, subfolder=subfolder, **load) + if lazy: + state = LazyNumpyStateDict(module) + else: + state = {} + for key, value in module.state_dict().items(): + if ".attentions." in key: + for old, new in legacy.items(): + key = key.replace(old, new) + state[key] = value.detach().cpu().numpy() + del module + consumed = set() + trainable, non_trainable = split_model_weights(component) + for keras_weight, _ in tqdm( + trainable + non_trainable, desc=f"Transferring {subfolder} weights to Keras" + ): + key = "/".join(keras_weight.path.split("/")[-2:]) + for old, new in WEIGHT_NAME_MAPPING.items(): + key = key.replace(old, new) + consumed.add(key) + if key not in state: + raise WeightMappingError(keras_weight.path, key) + torch_weight = state[key] + if not compare_keras_torch_names( + keras_weight.path, keras_weight, key, torch_weight + ): + raise WeightShapeMismatchError( + keras_weight.path, keras_weight.shape, key, torch_weight.shape + ) + if key.startswith(("time_embedding.", "add_embedding.")) and ( + keras_weight.ndim == 2 + ): + # transfer_weights treats any "embedding" 2D weight as a lookup + # table; the timestep / added-conditioning MLPs' linears are dense + # kernels + keras_weight.assign(np.transpose(torch_weight)) + continue + transfer_weights(key, keras_weight, torch_weight) + unused = sorted(set(state) - consumed) + if unused: + raise ValueError( + f"{type(component).__name__}: {len(unused)} checkpoint tensors " + f"unused, e.g. {unused[:3]}." + ) + del state + gc.collect() + + if config.text_config is not None: + text_encoder = CLIPTextModel.from_pretrained( + repo, subfolder="text_encoder", **fp32 + ) + transfer_text_encoder(model.text_encoder, text_encoder) + del text_encoder + gc.collect() + + text_encoder_2 = CLIPTextModelWithProjection.from_pretrained( + repo, subfolder="text_encoder_2", **fp32 + ) + transfer_text_encoder(model.text_encoder_2, text_encoder_2) + del text_encoder_2 + gc.collect() + return model, config + + +if __name__ == "__main__": + import gc + import os + + import keras + + OUT_DIR = os.environ.get("ZM_OUT_DIR", "C:/Users/gites/Desktop/code/v1_weights") + os.makedirs(OUT_DIR, exist_ok=True) + MAX_SHARD_GB = 5.0 + token = os.environ.get("HF_TOKEN") + selected = [v for v in os.environ.get("ZM_VARIANTS", "").split(",") if v] + sources = { + variant: source + for variant, source in STABLE_DIFFUSION_XL_SOURCES.items() + if not selected or variant in selected + } + + for variant, source in sources.items(): + print(f"\n{'=' * 60}\nConverting: {variant} <- {source}\n{'=' * 60}") + model, config = transfer_stable_diffusion_xl(source, token=token) + + n_bytes = sum( + int(np.prod(w.shape)) * np.dtype(w.dtype).itemsize for w in model.weights + ) + stem = os.path.join(OUT_DIR, variant.replace("-", "_").replace(".", "_")) + if n_bytes > MAX_SHARD_GB * 1024**3: + out = f"{stem}.weights.json" + model.save_weights(out, max_shard_size=MAX_SHARD_GB) + else: + out = f"{stem}.weights.h5" + model.save_weights(out) + print(f" saved -> {out} ({n_bytes / 1024**3:.2f} GB)") + + del model + keras.backend.clear_session() + gc.collect() diff --git a/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_config.py b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_config.py new file mode 100644 index 00000000..0ed30070 --- /dev/null +++ b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_config.py @@ -0,0 +1,215 @@ +from zeromodels.models.clip.clip_config import CLIPTextConfig +from zeromodels.models.stable_diffusion.stable_diffusion_config import ( + AutoencoderKLConfig, + StableDiffusionConfig, + StableDiffusionTextConfig, + UNet2DConditionConfig, +) + + +class StableDiffusionXLUNetConfig(UNet2DConditionConfig): + r"""The Stable Diffusion XL denoiser: a [`UNet2DConditionConfig`] with the SDXL + layout. + + Same fields as [`UNet2DConditionConfig`]; the defaults are the SDXL base UNet + (2.57B parameters): three levels ((320, 640, 1280) wide, no attention at the + first), a 2048-d text context (the two text encoders concatenated), one head + count per level ((5, 10, 20), a 64-wide head everywhere), (1, 2, 10) transformer + blocks per Transformer2D, a linear token projection and the ``text_time`` + micro-conditioning (the pooled 1280-d text embedding plus six size / crop ids + embedded to 256 each, 2816 in all). ``sample_size`` is 128 (1024px); SDXL-Turbo + ships 64 (512px).""" + + sample_size: int = 128 + down_block_types: tuple = ( + "DownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + ) + up_block_types: tuple = ("CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D") + block_out_channels: tuple = (320, 640, 1280) + cross_attention_dim: int = 2048 + num_attention_heads: int | tuple = (5, 10, 20) + use_linear_projection: bool = True + transformer_layers_per_block: int | tuple = (1, 2, 10) + addition_embed_type: str | None = "text_time" + addition_time_embed_dim: int = 256 + projection_class_embeddings_input_dim: int | None = 2816 + num_time_ids: int = 6 + + +class StableDiffusionXLVAEConfig(AutoencoderKLConfig): + r"""The Stable Diffusion XL VAE: an [`AutoencoderKLConfig`] built for 1024px with + the SDXL latent scaling (0.13025) and ``force_upcast`` on (this VAE overflows in + float16, so it is always built in float32).""" + + sample_size: int = 1024 + scaling_factor: float = 0.13025 + force_upcast: bool = True + + +class StableDiffusionXLTextConfig2(CLIPTextConfig): + r"""The second Stable Diffusion XL text tower: OpenCLIP ViT-bigG/14's text encoder. + + Same fields as [`CLIPTextConfig`], defaulted to the ViT-bigG/14 sizes (1280 wide, + 20 heads, 32 layers), plus the joint-space projection whose pooled output is the + UNet's ``text_embeds`` conditioning. + + Args: + projection_dim (`int`, *optional*, defaults to 1280): + Width of the ``text_projection`` applied to the pooled (EOT) state. + hidden_act (`str`, *optional*, defaults to `"gelu"`): + Activation of this tower (the first tower uses the top-level + ``hidden_act``, ``quick_gelu``).""" + + hidden_dim: int = 1280 + num_heads: int = 20 + num_layers: int = 32 + projection_dim: int = 1280 + hidden_act: str = "gelu" + + +class StableDiffusionXLConfig(StableDiffusionConfig): + r"""Configuration for [`StableDiffusionXLModel`], the hosted SDXL weights container. + + The [`StableDiffusionConfig`] shape (nested sub-configs plus the scheduler and + token ids; flat constructor) with a fourth component: SDXL conditions on two + text encoders, CLIP ViT-L/14's (``text_config``, the SD 1.x tower) and OpenCLIP + ViT-bigG/14's (``text_config_2``, with a 1280-d projection), whose penultimate + hidden states are concatenated into the UNet's 2048-d context while the second + tower's projected pooled state feeds the ``text_time`` micro-conditioning. The + flat constructor prefixes the second tower's fields ``text_2_`` + (``text_2_hidden_dim``, ...). + + Args: + unet_config ([`StableDiffusionXLUNetConfig`], *optional*): The denoiser. + vae_config ([`StableDiffusionXLVAEConfig`], *optional*): The VAE. + text_config ([`StableDiffusionTextConfig`], *optional*): The CLIP ViT-L/14 + text encoder. + text_config_2 ([`StableDiffusionXLTextConfig2`], *optional*): The OpenCLIP + ViT-bigG/14 text encoder. + pad_token_id_2 (`int`, *optional*, defaults to 0): + The second tokenizer pads with ``!`` (id 0) where the first pads with + ``<|endoftext|>``; the model derives the second tower's ids from the + first tokenizer's output and this id. + force_zeros_for_empty_prompt (`bool`, *optional*, defaults to True): + Use zero embeddings, rather than the encoded empty prompt, as the + unconditional branch of classifier-free guidance when no negative + prompt is given (the SDXL repos' ``model_index.json`` setting). + requires_aesthetics_score (`bool`, *optional*, defaults to False): + The refiner's micro-conditioning: an aesthetic score as the fifth time + id instead of the target size (``num_time_ids`` 5). + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_xl import StableDiffusionXLConfig, StableDiffusionXLModel + + >>> config = StableDiffusionXLConfig() + >>> model = StableDiffusionXLModel(config) # random weights, SDXL shapes + >>> model.config.unet_config.transformer_layers_per_block + (1, 2, 10) + ```""" + + model_type = "stable_diffusion_xl" + + sub_configs = { + "unet_config": StableDiffusionXLUNetConfig, + "vae_config": StableDiffusionXLVAEConfig, + "text_config": StableDiffusionTextConfig, + "text_config_2": StableDiffusionXLTextConfig2, + } + sub_config_prefixes = { + "unet_config": "unet_", + "vae_config": "vae_", + "text_config": "text_", + "text_config_2": "text_2_", + } + group_extras = {"text_config": ("vocab_size", "max_seq_len")} + + unet_config: StableDiffusionXLUNetConfig | dict | None = None + vae_config: StableDiffusionXLVAEConfig | dict | None = None + text_config_2: StableDiffusionXLTextConfig2 | dict | None = None + pad_token_id_2: int = 0 + force_zeros_for_empty_prompt: bool = True + requires_aesthetics_score: bool = False + + +class StableDiffusionXLRefinerUNetConfig(StableDiffusionXLUNetConfig): + r"""The Stable Diffusion XL refiner's denoiser: a [`StableDiffusionXLUNetConfig`] + with the refiner layout. + + Four levels ((384, 768, 1536, 1536) wide, attention at the middle two), 4 + transformer blocks per Transformer2D, (6, 12, 24, 24) heads, a 1280-d context + (the OpenCLIP ViT-bigG/14 tower alone) and five micro-conditioning values (size, + crop and the aesthetic score, 2560 with the pooled embedding). 2.3B parameters.""" + + down_block_types: tuple = ( + "DownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "DownBlock2D", + ) + up_block_types: tuple = ( + "UpBlock2D", + "CrossAttnUpBlock2D", + "CrossAttnUpBlock2D", + "UpBlock2D", + ) + block_out_channels: tuple = (384, 768, 1536, 1536) + cross_attention_dim: int = 1280 + num_attention_heads: int | tuple = (6, 12, 24, 24) + transformer_layers_per_block: int | tuple = 4 + projection_class_embeddings_input_dim: int | None = 2560 + num_time_ids: int = 5 + + +class StableDiffusionXLRefinerConfig(StableDiffusionXLConfig): + r"""Configuration for [`StableDiffusionXLRefinerModel`], the hosted SDXL refiner + weights container. + + The [`StableDiffusionXLConfig`] shape with the refiner's parts: no first text + tower (``text_config`` is ``None``; the OpenCLIP ViT-bigG/14 ``text_config_2`` + alone conditions the UNet, 1280-d), the [`StableDiffusionXLRefinerUNetConfig`] + denoiser and the aesthetic-score micro-conditioning + (``requires_aesthetics_score``); the empty negative prompt is encoded rather than + zeroed (``force_zeros_for_empty_prompt`` off), as the refiner repos declare. + + Args: + unet_config ([`StableDiffusionXLRefinerUNetConfig`], *optional*): The denoiser. + vae_config ([`StableDiffusionXLVAEConfig`], *optional*): The VAE. + text_config (*optional*): Absent (``None``): the refiner has no CLIP ViT-L/14 + tower. + text_config_2 ([`StableDiffusionXLTextConfig2`], *optional*): The OpenCLIP + ViT-bigG/14 text encoder. + requires_aesthetics_score (`bool`, *optional*, defaults to True): + Condition on an aesthetic score (the fifth time id) instead of a target + size. + force_zeros_for_empty_prompt (`bool`, *optional*, defaults to False): + Encode the empty prompt for the unconditional branch. + + Examples: + + ```python + >>> from zeromodels.models.stable_diffusion_xl import StableDiffusionXLRefinerConfig, StableDiffusionXLRefinerModel + + >>> config = StableDiffusionXLRefinerConfig() + >>> model = StableDiffusionXLRefinerModel(config) # random weights, refiner shapes + >>> model.config.unet_config.block_out_channels + (384, 768, 1536, 1536) + ```""" + + model_type = "stable_diffusion_xl_refiner" + + sub_configs = { + "unet_config": StableDiffusionXLRefinerUNetConfig, + "vae_config": StableDiffusionXLVAEConfig, + "text_config": StableDiffusionTextConfig, + "text_config_2": StableDiffusionXLTextConfig2, + } + optional_sub_configs = ("text_config",) + + unet_config: StableDiffusionXLRefinerUNetConfig | dict | None = None + text_config: StableDiffusionTextConfig | dict | None = None + requires_aesthetics_score: bool = True + force_zeros_for_empty_prompt: bool = False diff --git a/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_model.py b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_model.py new file mode 100644 index 00000000..ccb94a9d --- /dev/null +++ b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_model.py @@ -0,0 +1,496 @@ +import keras +from keras import layers, ops + +from zeromodels.base.base_scheduler import EulerDiscreteScheduler +from zeromodels.models.clip.clip_layers import CLIPTextModelEmbedding +from zeromodels.models.clip.clip_model import residual_attention_block +from zeromodels.models.stable_diffusion.stable_diffusion_model import ( + AutoencoderKL, + StableDiffusionModel, + StableDiffusionTextToImage, + UNet2DConditionModel, +) + +from .stable_diffusion_xl_config import ( + StableDiffusionXLConfig, + StableDiffusionXLRefinerConfig, +) + +# The container and the task build the identical graph, so both load the one +# hosted repo (whose zm_config.json names the container). +STABLE_DIFFUSION_XL_HUB_SIBLINGS = frozenset( + {"StableDiffusionXLModel", "StableDiffusionXLTextToImage"} +) +STABLE_DIFFUSION_XL_REFINER_HUB_SIBLINGS = frozenset( + {"StableDiffusionXLRefinerModel", "StableDiffusionXLRefinerImageToImage"} +) + + +def build_text_encoder( + config, hidden_act, layer_norm_eps, projection_dim=None, name="text_encoder" +): + """Functional CLIP text tower (the ``clip`` family's layers and leaf names) + that also exposes the penultimate hidden state SDXL conditions on. + + Outputs ``penultimate_hidden_state`` (the input of the last transformer block, + no final LayerNorm: diffusers' ``hidden_states[-2]``), the usual + ``last_hidden_state`` / ``pooler_output`` (EOT position, argmax of the ids) and, + when ``projection_dim`` is set, ``text_embeds``, the pooled state through the + bias-free ``text_projection`` (the second SDXL tower). Built under a + ``keras.name_scope(name)`` so the two towers' identically named layers keep + distinct weight paths inside the container. + """ + length = config.max_seq_len + with keras.name_scope(name): + # int32 inputs: under a float16 load policy a float input would round the + # ids (49407 is not representable in float16) + token_ids = layers.Input(shape=(length,), dtype="int32", name="token_ids") + padding_mask = layers.Input(shape=(length,), dtype="int32", name="padding_mask") + x = CLIPTextModelEmbedding( + vocab_size=config.vocab_size, + max_seq_len=length, + embed_dim=config.hidden_dim, + name="text_model_embedding", + )(token_ids) + + causal_mask = ops.cast(ops.triu(ops.ones((length, length)), k=1), "float32") + causal_mask = causal_mask * (-1e8) + key_mask = ops.reshape(ops.cast(padding_mask, "float32"), (-1, 1, 1, length)) + key_mask = (1.0 - ops.repeat(key_mask, length, axis=2)) * (-1e8) + + penultimate = None + for i in range(config.num_layers): + if i == config.num_layers - 1: + penultimate = x + x = residual_attention_block( + x, + proj_dim=config.hidden_dim, + num_heads=config.num_heads, + layer_name_prefix="text_model_encoder", + layer_idx=i, + causal_attention_mask=causal_mask, + attention_mask=key_mask, + mlp_ratio=config.mlp_ratio, + hidden_act=hidden_act, + layer_norm_eps=layer_norm_eps, + ) + last_hidden_state = layers.LayerNormalization( + epsilon=layer_norm_eps, name="text_model_layernorm" + )(x) + one_hot = ops.one_hot( + ops.argmax(token_ids, axis=-1), length, dtype=last_hidden_state.dtype + ) + pooler_output = ops.einsum("bi,bij->bj", one_hot, last_hidden_state) + + outputs = { + "penultimate_hidden_state": penultimate, + "last_hidden_state": last_hidden_state, + "pooler_output": pooler_output, + } + if projection_dim is not None: + # the Dense runs on a (B, 1, D) view, as CLIPTextEmbed does, so the + # kernel maps to the checkpoint's text_projection.weight unchanged + projected = layers.Dense( + projection_dim, use_bias=False, name="text_projection" + )(ops.expand_dims(pooler_output, axis=1)) + outputs["text_embeds"] = ops.squeeze(projected, axis=1) + return keras.Model( + inputs={"token_ids": token_ids, "padding_mask": padding_mask}, + outputs=outputs, + name=name, + ) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusionXLModel(StableDiffusionModel): + """Stable Diffusion XL weights container: the UNet, the VAE and the two text + encoders (CLIP ViT-L/14 and OpenCLIP ViT-bigG/14) as one functional + ``keras.Model``. + + The SD 1.x container (:class:`StableDiffusionModel`) with the SDXL + configuration and a fourth component: four disconnected sub-graphs, + ``{"sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids"} + -> "noise_pred"`` (the UNet with its ``text_time`` micro-conditioning), + ``{"image", "latent"} -> "moments" / "image"`` (the VAE, built in float32 + whatever the load dtype, ``force_upcast``) and ``{"token_ids", "token_ids_2", + "padding_mask"} -> "prompt_embeds" / "pooled_prompt_embeds"`` (the two text + towers: their penultimate hidden states concatenated to the UNet's 2048-d + context, and the second tower's projected pooled state). The components are + ``.unet`` / ``.vae`` / ``.text_encoder`` / ``.text_encoder_2``. Hosted as one + ``zm_config.json`` + float16 weights (the checkpoints' native precision) per + variant under ``zeromodels/stable-diffusion-xl-base-1.0`` and + ``zeromodels/sdxl-turbo``; on-the-fly ``hf:`` conversion is not supported for + diffusion models. + + Args: + **kwargs: The flat :class:`StableDiffusionXLConfig` fields (``unet_`` / + ``vae_`` / ``text_`` / ``text_2_`` prefixed), or a config positionally. + Pass ``unet_sample_size`` / ``vae_sample_size`` to build for another + resolution (the weights are resolution-independent). + """ + + config_class = StableDiffusionXLConfig + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_XL_HUB_SIBLINGS + + def __init__(self, name="StableDiffusionXLModel", **kwargs): + super().__init__(name=name, **kwargs) + + def build_components(self, config, data_format, channels_axis): + u, v, t, t2 = ( + config.unet_config, + config.vae_config, + config.text_config, + config.text_config_2, + ) + unet = UNet2DConditionModel( + u, data_format=data_format, channels_axis=channels_axis + ) + vae = AutoencoderKL( + **v.constructor_kwargs(), + data_format=data_format, + channels_axis=channels_axis, + ) + components = {"unet": unet, "vae": vae} + if t is not None: # the refiner has no CLIP ViT-L/14 tower + components["text_encoder"] = build_text_encoder( + t, config.hidden_act, config.layer_norm_eps, name="text_encoder" + ) + components["text_encoder_2"] = build_text_encoder( + t2, + t2.hidden_act, + config.layer_norm_eps, + projection_dim=t2.projection_dim, + name="text_encoder_2", + ) + return components + + def build_graph(self, config, components, data_format): + u, length = config.unet_config, config.text_config_2.max_seq_len + unet, vae = components["unet"], components["vae"] + text_encoder = components.get("text_encoder") + text_encoder_2 = components["text_encoder_2"] + inputs = self.unet_vae_inputs(config, components, data_format) + pooled_dim = ( + u.projection_class_embeddings_input_dim + - u.num_time_ids * u.addition_time_embed_dim + ) + inputs["text_embeds"] = layers.Input(shape=(pooled_dim,), name="text_embeds") + inputs["time_ids"] = layers.Input(shape=(u.num_time_ids,), name="time_ids") + if text_encoder is not None: + inputs["token_ids"] = layers.Input( + shape=(length,), dtype="int32", name="token_ids" + ) + inputs["token_ids_2"] = layers.Input( + shape=(length,), dtype="int32", name="token_ids_2" + ) + inputs["padding_mask"] = layers.Input( + shape=(length,), dtype="int32", name="padding_mask" + ) + + noise_pred = unet( + { + key: inputs[key] + for key in ( + "sample", + "timestep", + "encoder_hidden_states", + "text_embeds", + "time_ids", + ) + } + )["sample"] + vae_out = vae({"image": inputs["image"], "latent": inputs["latent"]}) + text_out_2 = text_encoder_2( + { + "token_ids": inputs["token_ids_2"], + "padding_mask": inputs["padding_mask"], + } + ) + hidden_states = [text_out_2["penultimate_hidden_state"]] + if text_encoder is not None: + text_out = text_encoder( + { + "token_ids": inputs["token_ids"], + "padding_mask": inputs["padding_mask"], + } + ) + hidden_states.insert(0, text_out["penultimate_hidden_state"]) + prompt_embeds = ( + ops.concatenate(hidden_states, axis=-1) + if len(hidden_states) > 1 + else hidden_states[0] + ) + outputs = { + "noise_pred": noise_pred, + "moments": vae_out["moments"], + "image": vae_out["sample"], + "prompt_embeds": prompt_embeds, + "pooled_prompt_embeds": text_out_2["text_embeds"], + } + return inputs, outputs + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusionXLTextToImage(StableDiffusionTextToImage, StableDiffusionXLModel): + """Text-to-image Stable Diffusion XL, pure Keras 3 and cross-backend (the + diffusers ``StableDiffusionXLPipeline`` in this library's task-class form). + + :class:`StableDiffusionTextToImage`'s ``generate`` (``BaseDiffusion``) over the + :class:`StableDiffusionXLModel` graph. The hooks follow the SDXL recipe: the + prompt goes through both text towers (the second tower's ``!``-padded ids are + derived from the first tokenizer's output), their penultimate hidden states + are concatenated into the UNet context and the second tower's projected pooled + state becomes the ``text_embeds`` micro-conditioning, next to the ``time_ids`` + built from ``original_size`` / ``crops_coords_top_left`` / ``target_size``. + With no negative prompt the unconditional branch of classifier-free guidance + is zero embeddings (``force_zeros_for_empty_prompt``), as in the reference. + The sampler is the repo's ``scheduler_config`` (Euler with ``leading`` spacing + for the base model, ancestral Euler with ``trailing`` spacing for SDXL-Turbo):: + + sd = StableDiffusionXLTextToImage.from_weights("zeromodels/stable-diffusion-xl-base-1.0") + tokenizer = StableDiffusionXLTokenizer.from_weights("zeromodels/stable-diffusion-xl-base-1.0") + image = sd.generate(**tokenizer("a photo of a cat")) # (1, 1024, 1024, 3) uint8 + + Args: + scheduler: A :class:`BaseScheduler`; defaults to the config's + ``scheduler_config``, else Euler with ``leading`` spacing. + **kwargs: The flat :class:`StableDiffusionXLConfig` fields. + """ + + config_class = StableDiffusionXLConfig + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_XL_HUB_SIBLINGS + # diffusers' StableDiffusionXLPipeline defaults; a repo's generate_args override + generate_args = {"num_inference_steps": 50, "guidance_scale": 5.0} + + def __init__(self, scheduler=None, name="StableDiffusionXLTextToImage", **kwargs): + super().__init__(scheduler=scheduler, name=name, **kwargs) + + def default_scheduler(self): + return EulerDiscreteScheduler(timestep_spacing="leading", steps_offset=1) + + def generate( + self, + input_ids, + attention_mask=None, + negative_input_ids=None, + num_inference_steps=None, + guidance_scale=None, + seed=None, + latents=None, + image=None, + strength=None, + denoising_start=None, + denoising_end=None, + output_type="image", + original_size=None, + crops_coords_top_left=(0, 0), + target_size=None, + aesthetic_score=6.0, + negative_original_size=None, + negative_crops_coords_top_left=None, + negative_target_size=None, + negative_aesthetic_score=2.5, + ): + """Generate images from tokenized prompts, or refine an ``image`` / latent + (see ``BaseDiffusion.generate`` for the shared arguments). + + The SDXL micro-conditioning arguments are ``(height, width)`` pairs in + pixels: ``original_size`` and ``target_size`` default to the generated + image's size, ``crops_coords_top_left`` to ``(0, 0)``; the ``negative_`` + variants condition the negative branch and default to the positive values. + ``aesthetic_score`` / ``negative_aesthetic_score`` (6.0 / 2.5) replace the + target size on the refiner (``requires_aesthetics_score``). For the SDXL + ensemble, run the base with ``denoising_end=0.8, output_type="latent"`` and + hand the latent to the refiner's ``generate(..., latents=..., + denoising_start=0.8)``. + """ + return super().generate( + input_ids, + attention_mask=attention_mask, + negative_input_ids=negative_input_ids, + num_inference_steps=num_inference_steps, + guidance_scale=guidance_scale, + seed=seed, + latents=latents, + image=image, + strength=strength, + denoising_start=denoising_start, + denoising_end=denoising_end, + output_type=output_type, + original_size=original_size, + crops_coords_top_left=crops_coords_top_left, + target_size=target_size, + aesthetic_score=aesthetic_score, + negative_original_size=negative_original_size, + negative_crops_coords_top_left=negative_crops_coords_top_left, + negative_target_size=negative_target_size, + negative_aesthetic_score=negative_aesthetic_score, + ) + + def image_size(self): + size = self.vae.sample_size + return tuple(size) if isinstance(size, (tuple, list)) else (size, size) + + def time_ids( + self, + batch, + original_size=None, + crops_coords_top_left=(0, 0), + target_size=None, + aesthetic_score=6.0, + ): + """The ``(batch, 6)`` float32 micro-conditioning row: original size, crop + offset and target size, each ``(height, width)``; ``(batch, 5)`` with the + aesthetic score in place of the target size when + ``requires_aesthetics_score`` (the refiner).""" + original_size = tuple(original_size or self.image_size()) + row = list(original_size) + list(crops_coords_top_left) + if self.config.requires_aesthetics_score: + row = row + [aesthetic_score] + else: + row = row + list(target_size or self.image_size()) + return ops.convert_to_tensor([row] * batch, dtype="float32") + + def unconditional_ids(self, batch): + cfg = self.config + length = cfg.text_config_2.max_seq_len + row = [cfg.bos_token_id, cfg.eos_token_id] + [cfg.pad_token_id] * (length - 2) + return ops.convert_to_tensor([row] * batch, dtype="int32") + + def prompt_mask(self, input_ids): + # the tokens up to and including the first <|endoftext|>; what follows is + # padding (the first tokenizer pads with <|endoftext|> itself) + is_eos = ops.cast(ops.equal(input_ids, self.config.eos_token_id), "int32") + return ops.equal(ops.cumsum(is_eos, axis=1) - is_eos, 0) + + def encode_prompt( + self, + input_ids, + attention_mask=None, + original_size=None, + crops_coords_top_left=(0, 0), + target_size=None, + aesthetic_score=6.0, + ): + input_ids = ops.cast(ops.convert_to_tensor(input_ids), "int32") + if attention_mask is None: + real = self.prompt_mask(input_ids) + else: + real = ops.greater( + ops.cast(ops.convert_to_tensor(attention_mask), "int32"), 0 + ) + # the second tokenizer pads with "!" (id 0) where the first pads with + # <|endoftext|>; the towers attend the padding (no attention mask), as the + # reference does + input_ids_2 = ops.where(real, input_ids, self.config.pad_token_id_2) + ones = ops.ones_like(input_ids) + out_2 = self.text_encoder_2({"token_ids": input_ids_2, "padding_mask": ones}) + hidden = out_2["penultimate_hidden_state"] + if self.config.text_config is not None: + out = self.text_encoder({"token_ids": input_ids, "padding_mask": ones}) + hidden = ops.concatenate([out["penultimate_hidden_state"], hidden], -1) + batch = int(input_ids.shape[0]) + return { + "encoder_hidden_states": hidden, + "text_embeds": out_2["text_embeds"], + "time_ids": self.time_ids( + batch, + original_size, + crops_coords_top_left, + target_size, + aesthetic_score, + ), + } + + def encode_negative_prompt(self, negative_input_ids, batch, **conditioning): + if ( + negative_input_ids is not None + or not self.config.force_zeros_for_empty_prompt + ): + return super().encode_negative_prompt( + negative_input_ids, batch, **conditioning + ) + # force_zeros_for_empty_prompt: zero embeddings stand in for the empty + # prompt; the micro-conditioning keeps its (negative) values + u, length = self.config.unet_config, self.config.text_config_2.max_seq_len + pooled_dim = ( + u.projection_class_embeddings_input_dim + - u.num_time_ids * u.addition_time_embed_dim + ) + return { + "encoder_hidden_states": ops.zeros( + (batch, length, u.cross_attention_dim), dtype="float32" + ), + "text_embeds": ops.zeros((batch, pooled_dim), dtype="float32"), + "time_ids": self.time_ids(batch, **conditioning), + } + + def predict_noise(self, latents, timesteps, embeddings): + return self.unet({"sample": latents, "timestep": timesteps, **embeddings})[ + "sample" + ] + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusionXLRefinerModel(StableDiffusionXLModel): + """Stable Diffusion XL refiner weights container: the refiner UNet, the VAE and + the OpenCLIP ViT-bigG/14 text encoder as one functional ``keras.Model``. + + :class:`StableDiffusionXLModel` with the refiner configuration + (:class:`StableDiffusionXLRefinerConfig`): three disconnected sub-graphs (no + CLIP ViT-L/14 tower, so the UNet's context is the second tower's 1280-d + penultimate state alone), the four-level refiner UNet and five + micro-conditioning ids (size, crop, aesthetic score). Components: + ``.unet`` / ``.vae`` / ``.text_encoder_2``. Hosted under + ``zeromodels/stable-diffusion-xl-refiner-1.0`` (float16, the VAE float32). + + Args: + **kwargs: The flat :class:`StableDiffusionXLRefinerConfig` fields, or a + config positionally. + """ + + config_class = StableDiffusionXLRefinerConfig + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_XL_REFINER_HUB_SIBLINGS + + def __init__(self, name="StableDiffusionXLRefinerModel", **kwargs): + super().__init__(name=name, **kwargs) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusionXLRefinerImageToImage( + StableDiffusionXLTextToImage, StableDiffusionXLRefinerModel +): + """Image-to-image Stable Diffusion XL refiner, pure Keras 3 and cross-backend + (diffusers' ``StableDiffusionXLImg2ImgPipeline`` with the refiner weights). + + :class:`StableDiffusionXLTextToImage`'s hooks over the + :class:`StableDiffusionXLRefinerModel` graph, the prompt encoded by the + OpenCLIP ViT-bigG/14 tower alone and the aesthetic score + (``aesthetic_score`` 6.0 / ``negative_aesthetic_score`` 2.5) as the fifth + micro-conditioning id. ``generate`` refines an ``image`` (noised to + ``strength``, 0.3 by default) or, as the second half of the SDXL ensemble of + experts, a latent the base model left partially denoised + (``denoising_start`` at the base's ``denoising_end``):: + + base = StableDiffusionXLTextToImage.from_weights("zeromodels/stable-diffusion-xl-base-1.0") + refiner = StableDiffusionXLRefinerImageToImage.from_weights("zeromodels/stable-diffusion-xl-refiner-1.0") + tokenizer = StableDiffusionXLTokenizer.from_weights("zeromodels/stable-diffusion-xl-base-1.0") + inputs = tokenizer("a photo of a cat") + latent = base.generate(**inputs, denoising_end=0.8, output_type="latent") + image = refiner.generate(**inputs, latents=latent, denoising_start=0.8) + + Args: + scheduler: A :class:`BaseScheduler`; defaults to the config's + ``scheduler_config``, else Euler with ``leading`` spacing. + **kwargs: The flat :class:`StableDiffusionXLRefinerConfig` fields. + """ + + config_class = StableDiffusionXLRefinerConfig + HUB_REPO_SIBLINGS = STABLE_DIFFUSION_XL_REFINER_HUB_SIBLINGS + # diffusers' StableDiffusionXLImg2ImgPipeline defaults + generate_args = {"num_inference_steps": 50, "guidance_scale": 5.0, "strength": 0.3} + + def __init__( + self, scheduler=None, name="StableDiffusionXLRefinerImageToImage", **kwargs + ): + super().__init__(scheduler=scheduler, name=name, **kwargs) diff --git a/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_tokenizer.py b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_tokenizer.py new file mode 100644 index 00000000..88046a65 --- /dev/null +++ b/zeromodels/models/stable_diffusion_xl/stable_diffusion_xl_tokenizer.py @@ -0,0 +1,42 @@ +import keras + +from zeromodels.models.stable_diffusion.stable_diffusion_tokenizer import ( + StableDiffusionTokenizer, +) + + +@keras.saving.register_keras_serializable(package="zeromodels") +class StableDiffusionXLTokenizer(StableDiffusionTokenizer): + """Stable Diffusion XL text tokenizer: the CLIP ViT-L/14 BPE tokenizer. + + SDXL tokenizes the prompt twice, once per text encoder, with the same + byte-pair encoding and ``<|startoftext|>`` / ``<|endoftext|>`` framing: the + two tokenizers differ only in the pad token (``<|endoftext|>`` for the CLIP + ViT-L/14 tower, ``!`` for the OpenCLIP ViT-bigG/14 one). This is the first; + :class:`StableDiffusionXLTextToImage` derives the second tower's ids from the + returned ``attention_mask``, so one call feeds ``generate``. Loads by repo id + like the model: ``from_weights("zeromodels/stable-diffusion-xl-base-1.0")``. + + Args: + hf_id: Hosted repo to read ``tokenizer.json`` from. Required unless + ``tokenizer_file`` is given (no default repo). + tokenizer_file: Explicit ``tokenizer.json`` path (overrides ``hf_id``). + max_seq_len: Padded / truncated length (default 77). + pad_token: Pad token string (``<|endoftext|>``). + """ + + def __init__( + self, + hf_id=None, + tokenizer_file=None, + max_seq_len=77, + pad_token="<|endoftext|>", + **kwargs, + ): + super().__init__( + hf_id=hf_id, + tokenizer_file=tokenizer_file, + max_seq_len=max_seq_len, + pad_token=pad_token, + **kwargs, + )