[diffusion] model: support qwen-image-2.1 (#39983)
Co-authored-by: Mick Qian <mickqian@users.noreply.github.com> Co-authored-by: BBuf <1182563586@qq.com>
This commit is contained in:
@@ -0,0 +1,547 @@
|
||||
---
|
||||
title: Qwen-Image 2.1
|
||||
description: "Run Qwen-Image 2.1 text-to-image and image-conditioned generation with SGLang Diffusion."
|
||||
tag: NEW
|
||||
---
|
||||
|
||||
import { DiffusionModelTags } from '/src/snippets/diffusion/model-tags.jsx';
|
||||
import { Deployment } from '/src/snippets/_deployment.jsx';
|
||||
import { config } from '/src/snippets/configs/Qwen/qwen-image-2.1.jsx';
|
||||
|
||||
<DiffusionModelTags tags={["RGBA image", "text-to-image", "image editing", "multi-image references", "block-causal attention"]} />
|
||||
|
||||
## 1. Quick start
|
||||
|
||||
Install the runtime dependencies with `uv pip install "sglang[diffusion]" --prerelease=allow`,
|
||||
then install this integration from its source checkout with
|
||||
`uv pip install -e "python[diffusion]"`. Use an authorized checkpoint directory in
|
||||
place of `/models/qwen-image-2.1`. The recipes below target NVIDIA CUDA on Linux;
|
||||
the hardware picker selects a tested single-GPU recipe for the full checkpoint.
|
||||
|
||||
<Deployment config={config} />
|
||||
|
||||
Use **Setup** to select text-to-image, single-image editing, or multi-image
|
||||
editing. **Server** controls placement, attention, encoder scheduling, VAE
|
||||
tiling, and graph execution. **Request** controls the background, resolution,
|
||||
steps, and output count. Set reference PNG paths under **Variables**; edits
|
||||
upload files from the machine running cURL, so they need not exist on the server.
|
||||
|
||||
Hardware selection applies the recommended placement for that GPU. H200,
|
||||
B200, and RTX PRO 6000 96GB keep weights resident; RTX 5090 and RTX 4090 use
|
||||
offload to fit the full pipeline.
|
||||
Custom two- and four-GPU topologies and unverified feature combinations remain selectable and are labeled
|
||||
**Unverified**. Invalid topology combinations disable Copy. This integration
|
||||
currently uses the Python/source command; no published Docker image is verified.
|
||||
|
||||
Both request modes return base64 PNGs. To save all returned images, append
|
||||
`> response.json` to the request command, then run:
|
||||
|
||||
```bash Command
|
||||
python - <<'PY'
|
||||
import base64
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
for i, item in enumerate(json.loads(Path("response.json").read_text())["data"]):
|
||||
Path(f"output-{i}.png").write_bytes(base64.b64decode(item["b64_json"]))
|
||||
PY
|
||||
```
|
||||
|
||||
### Platform measurements
|
||||
|
||||
The following four-platform comparison and the fusion measurements below precede
|
||||
the training-template and VAE normalization corrections in `c2a31b2693c`;
|
||||
their output comparisons should not be treated as baselines for that revision.
|
||||
The separate RTX PRO 6000 measurement uses the corrected implementation.
|
||||
|
||||
| GPU | Recommended placement / attention | Generation median | Single edit | Peak device memory |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| H200 141GB | Resident / FlashAttention | Functional verification only | Passed | Not measured in this comparison |
|
||||
| B200 192GB | Resident / FlashAttention | 3.44 s | 3.84 s | 40.1 GiB |
|
||||
| RTX 5090 32GB | DiT layerwise offload / SDPA | 14.30 s | 16.84 s | 26.9 GiB |
|
||||
| RTX 4090 24GB | DiT layerwise + encoder CPU offload / FlashAttention | 24.60 s | 25.59 s | 21.4 GiB |
|
||||
|
||||
The recommendations compare exact attention backends and memory placement on
|
||||
one GPU per platform. Each run warms up with one 512px, 4-step request, then
|
||||
measures three 1024px, 40-step generations, one single-image edit, and one
|
||||
transparent generation. All use seed 42, CFG 1, eager execution, full-image VAE
|
||||
decoding, and PNG output. Generation latency is the median of three sequential
|
||||
HTTP requests; editing is one request. Times include encoding and PNG response
|
||||
serialization, but exclude server startup. Device memory is the highest sampled
|
||||
`nvidia-smi` usage across loading and requests, sampled every 0.5 seconds.
|
||||
|
||||
Measured on 2026-09-16 with source revision `128ae46cc`, PyTorch 2.13.0+cu130,
|
||||
Transformers 5.12.1, and Diffusers 0.37.0. SGLang's native encoder uses the
|
||||
Transformers 4.57.3 numerical semantics described below. The RTX 5090 runs used
|
||||
a 50 GiB process-group memory limit on a roughly 60 GiB host; this is a tested
|
||||
budget, not a minimum host-memory requirement.
|
||||
|
||||
B200 FlashAttention was faster than SDPA in this comparison (3.44 vs 3.70 s).
|
||||
On RTX 5090, both commands used Torch SDPA: this runtime falls back to SDPA
|
||||
when `--attention-backend fa` is selected on SM120. The measured 14.30 s
|
||||
(explicit SDPA) and 14.39 s (FA selection with SDPA fallback) therefore do not
|
||||
compare different backends. The picker defaults to SDPA and rejects Ring with
|
||||
either selection on RTX 5090. Keeping eight DiT layers resident
|
||||
did not improve the RTX 5090 generation median, so that flag is omitted.
|
||||
On RTX 4090, DiT offload alone passed generation but ran out of memory during
|
||||
editing. The recommended command also sets `--text-encoder-cpu-offload true`;
|
||||
this complete recipe passed generation, editing, and transparent PNG output.
|
||||
|
||||
These are measurements of this small workload, not universal latency or image
|
||||
quality guarantees. Different prompts, reference sizes, batching, and software
|
||||
versions can change memory use and latency. Multi-reference and batched request
|
||||
recipes retain their separate H200 verification scope in the picker.
|
||||
|
||||
### RTX PRO 6000 Blackwell 96GB
|
||||
|
||||
The recommended single-GPU command keeps all weights resident and selects Torch
|
||||
SDPA. This is the 96GB Blackwell Server Edition (SM120). This runtime also maps
|
||||
`--attention-backend fa` to SDPA on this GPU; Ring therefore requires another
|
||||
supported backend and is rejected with either selection in the picker.
|
||||
|
||||
Source revision `1eab5de5990` was measured on 2026-09-18:
|
||||
|
||||
| Placement | Generation median | Edit median | Peak device memory |
|
||||
| --- | --- | --- | --- |
|
||||
| Resident (recommended) | 8.23 s | 9.85 s | 40.1 GiB |
|
||||
| DiT layerwise offload | 10.28 s | 10.66 s | 26.1 GiB |
|
||||
|
||||
Both runs used PyTorch 2.13.0+cu130, Transformers 5.12.1, Diffusers 0.37.0,
|
||||
native precision, eager execution, and full-image VAE decoding.
|
||||
After two 1024px/40-step warmups, each measured five generations and three edits
|
||||
at that same resolution and step count, with seed 42, CFG 1, and CPU noise
|
||||
generation. HTTP latency includes PNG serialization and excludes server startup;
|
||||
device memory was sampled every 0.5 seconds across startup and requests.
|
||||
|
||||
Transparent generation and two repeated edits of the same transparent input passed
|
||||
with both placements, retaining alpha values from 0 to 255. Repeated requests
|
||||
and corresponding outputs across placements produced identical RGBA pixels for
|
||||
this workload. Quantized checkpoints and multi-GPU recipes on RTX PRO 6000 remain
|
||||
unverified.
|
||||
|
||||
### Lossless RoPE fusion
|
||||
|
||||
The native DiT fuses the float conversion, complex rotary multiplication, and
|
||||
output cast on supported CUDA tensors. Its first eager call checks exact
|
||||
agreement with the original PyTorch operation; a mismatch disables the fusion.
|
||||
No additional command flag is needed.
|
||||
|
||||
A separate comparison on 2026-09-17 used native revision `6b190085c48` as the
|
||||
baseline and `63ed20bbedb` with the fusion. Both used the software versions
|
||||
listed above, full-image VAE decode, eager execution, and the recommended
|
||||
placement and attention backend for each GPU:
|
||||
|
||||
- B200: generation **3.42 → 3.27 s** (4.5% lower latency), editing
|
||||
**4.03 → 3.89 s** (3.4% lower).
|
||||
- RTX 5090: generation **14.49 → 14.20 s** (2.0% lower), editing
|
||||
**16.97 → 16.68 s** (1.7% lower).
|
||||
|
||||
Each GPU ran four fresh servers in optimized/baseline/baseline/optimized order.
|
||||
Each startup used two full-size warmups followed by five generations and three
|
||||
edits. The medians pool 10 generations and six edits per variant, all at
|
||||
1024px, 40 steps, seed 42, CFG 1, CPU noise generation, and one RGBA PNG per
|
||||
request. The workload generated a red teapot and edited the same reference
|
||||
image to blue. HTTP times include PNG serialization and exclude startup.
|
||||
All corresponding output pixels were identical between revisions on each GPU.
|
||||
These measurements cover this fixed workload; other prompts and configurations
|
||||
can have different gains.
|
||||
|
||||
### Lossless MLP and residual fusion
|
||||
|
||||
The native DiT also uses the shared BF16 SiLU-multiply and gated-residual
|
||||
kernels, preserving the eager operations' intermediate rounding. SiLU-multiply
|
||||
checks its first eager call and falls back on mismatch. These optimizations
|
||||
are automatic on supported CUDA inputs.
|
||||
|
||||
A second B200 comparison on 2026-09-17 used `f874eae18be` (already including
|
||||
the RoPE fusion) versus `a3d14531474`. With resident weights, FlashAttention,
|
||||
and the same four-startup protocol and workload above, generation decreased
|
||||
from **3.272 to 3.134 s** (4.23%) and editing from **3.886 to 3.762 s** (3.18%).
|
||||
All corresponding RGBA pixels were identical across the 10 generation and six
|
||||
editing samples per variant. These are additional gains over the RoPE baseline;
|
||||
this comparison does not establish the gain on other GPUs.
|
||||
|
||||
### Lossless Q/K normalization
|
||||
|
||||
Q/K RMSNorm fuses the input conversion and square, then the normalization,
|
||||
output cast, and weight multiply. It retains the original FP32 mean reduction
|
||||
with the same tensor shape, preserving the eager reduction order and
|
||||
cast-before-weight rounding. The native DiT verifies its first eager call and
|
||||
uses the original implementation if the outputs differ. No flag is needed.
|
||||
|
||||
A B200 comparison on 2026-09-17 used revision `4e5459e0eda` (including the
|
||||
RoPE, MLP, and residual fusions) versus `d9e1e5dac96`. With resident weights,
|
||||
FlashAttention, and the four-startup protocol above, generation decreased from
|
||||
**3.114 to 2.828 s** (9.17%) and editing from **3.742 to 3.450 s** (7.81%).
|
||||
Each variant has 10 generation and six editing measurements at 1024px,
|
||||
40 steps, seed 42, and CFG 1. Every corresponding RGBA pixel was identical.
|
||||
These gains apply to this fixed B200 workload; other GPUs were not measured
|
||||
in this comparison.
|
||||
|
||||
### Lossless LayerNorm modulation
|
||||
|
||||
The DiT fuses affine-free LayerNorm and `* (1 + scale)` while retaining the
|
||||
eager Welford reduction and BF16 rounding order. Scale-only modulation skips
|
||||
the shift addition, including its effect on signed zeros. The first eager call
|
||||
checks the fused result against the native path and falls back on a mismatch.
|
||||
|
||||
A B200 comparison on 2026-09-17 used `5bddbfca9b1` (including the preceding
|
||||
fusions) versus `162181ff0ec`. With resident weights, FlashAttention, eager
|
||||
execution, and the same four-startup protocol, generation decreased from
|
||||
**2.831 to 2.748 s** (2.92%) and editing from **3.436 to 3.358 s** (2.26%).
|
||||
Each variant has 10 generation and six editing measurements at 1024px,
|
||||
40 steps, seed 42, and CFG 1. Every corresponding RGBA pixel was identical.
|
||||
This comparison measures this B200 workload only.
|
||||
|
||||
## 2. Model capabilities
|
||||
|
||||
Qwen-Image 2.1 supports text-to-image generation and image-conditioned editing
|
||||
through one pipeline. Qwen3-VL encodes the instruction and reference images;
|
||||
a single-stream transformer inserts each reference image's latents into its
|
||||
corresponding position in that sequence. Block-causal attention keeps each
|
||||
image internally bidirectional while respecting the order of text and images.
|
||||
|
||||
For successive edits, send the previous output as the next request's reference
|
||||
image. Requests do not retain dialogue history. Conditional KV is reused across
|
||||
denoising steps within one request and released afterward; cross-request caching
|
||||
and incremental dialogue-history caching are not implemented.
|
||||
|
||||
Choose this pipeline for checkpoints declaring `QwenImage21Pipeline`,
|
||||
`QwenImage21Transformer2DModel`, and `AutoencoderKLQwenImage21`. The older
|
||||
Qwen-Image and Qwen-Image-Edit checkpoints use different components and latent
|
||||
packing. They cannot share this model's VAE or transformer weights. Text and
|
||||
condition-image activations use timestep zero, allowing their attention keys
|
||||
and values to be reused for the remaining denoising steps.
|
||||
|
||||
## 3. Checkpoint layout
|
||||
|
||||
The checkpoint directory must contain `model_index.json` and the `processor`,
|
||||
`text_encoder`, `transformer`, `vae`, and `scheduler` subdirectories. The
|
||||
processor must include the Qwen3-VL tokenizer assets. SGLang loads all three
|
||||
neural components natively. A separate tokenizer directory is not required.
|
||||
|
||||
The checkpoint's VAE uses RGBA input and output with 64-channel latents. PNG
|
||||
reference images retain their alpha channel; RGB inputs receive an opaque
|
||||
alpha channel. Save generated images as PNG to preserve transparency.
|
||||
|
||||
Text conditioning uses the last decoder layer's output before the final
|
||||
normalization, matching the reference implementation with Transformers
|
||||
4.57.3. Vision position interpolation also follows its BF16 rounding order.
|
||||
SGLang selects these native semantics explicitly, so keep the
|
||||
repository's installed dependencies instead of downgrading the entire runtime.
|
||||
The updated [Diffusers reference](https://github.com/huggingface/diffusers/pull/14804)
|
||||
also selects pre-normalization hidden states explicitly on newer Transformers.
|
||||
|
||||
Editing uses the training markers `<image1>`, `<image2>`, and so on. The vision
|
||||
encoder sees alpha composited over white, while the VAE receives the original
|
||||
RGBA pixels. Empty prompts become a space. The VAE normalizes features in
|
||||
FP32 before casting back to the activation dtype and compresses spatial
|
||||
dimensions by a factor of 16.
|
||||
|
||||
Use `--model-id Qwen-Image-2.1` when the checkpoint directory has a different
|
||||
name. The model ID is a routing identifier; it does not grant access to model
|
||||
weights. Keep checkpoint access credentials in your environment.
|
||||
|
||||
### Two-GPU end-to-end test
|
||||
|
||||
The `qwen_image21_t2i_tp2` case is temporarily disabled until the checkpoint is
|
||||
accessible to fork PR CI. Its configuration and pinned reference image are
|
||||
retained for re-enabling the test.
|
||||
|
||||
The case uses TP 2 with sequence
|
||||
parallelism disabled, 1024 × 1024 PNG output, 40 steps, CFG 1, and seed 42.
|
||||
It sends two consecutive requests and checks the model API and image consistency.
|
||||
This case does not enforce a latency baseline or run a component accuracy check.
|
||||
|
||||
### Transparent PNG output
|
||||
|
||||
Choose **Transparent / alpha** under Request to generate an isolated subject
|
||||
or preserve a transparent reference during editing. The picker adds the
|
||||
transparency instruction to the prompt and sets `output_format: "png"`.
|
||||
`background: "transparent"` alone only selects an output format; it does not
|
||||
remove the background or change model conditioning. JPEG cannot retain alpha.
|
||||
|
||||
The model predicts continuous alpha values, including partly transparent edges.
|
||||
No thresholding or background-removal postprocessing is applied. Transparent
|
||||
generation and transparent-input editing were compared against the reference
|
||||
at 1024 × 1024 and 40 steps; that check does not guarantee perfect cutouts for
|
||||
every prompt. Transparent generation and single-image editing also passed on
|
||||
the recommended one-H200 and one-RTX PRO 6000 servers at that resolution and
|
||||
step count, with one output per request.
|
||||
|
||||
## 4. Offline requests
|
||||
|
||||
### Text-to-image
|
||||
|
||||
```bash Command
|
||||
sglang generate \
|
||||
--model-path /models/qwen-image-2.1 \
|
||||
--model-id Qwen-Image-2.1 \
|
||||
--prompt "A capybara reading a book by candlelight" \
|
||||
--width 1024 --height 1024 \
|
||||
--num-inference-steps 40 --guidance-scale 1 \
|
||||
--seed 0 --save-output
|
||||
```
|
||||
|
||||
### Image-conditioned editing
|
||||
|
||||
```bash Command
|
||||
sglang generate \
|
||||
--model-path /models/qwen-image-2.1 \
|
||||
--model-id Qwen-Image-2.1 \
|
||||
--image-path /path/to/input.png \
|
||||
--prompt "Move the scene to a snowy mountain at sunrise" \
|
||||
--width 1024 --height 1024 \
|
||||
--num-inference-steps 40 --guidance-scale 1 \
|
||||
--seed 0 --save-output
|
||||
```
|
||||
|
||||
Height and width must be positive multiples of 32. Reference images preserve
|
||||
their aspect ratio and are resized to approximately the requested output area;
|
||||
the same resized image feeds the VLM and VAE. Image labels are deterministic
|
||||
(`Picture 1`, `Picture 2`, and so on). Multiple outputs receive independent
|
||||
noise seeds and independent prefix caches.
|
||||
|
||||
## 5. Runtime features
|
||||
|
||||
The API requires a text prompt; precomputed embeddings alone do not provide
|
||||
the image-token positions needed by this pipeline.
|
||||
|
||||
The default is 40 Euler flow-matching steps with CFG disabled. To use CFG,
|
||||
provide `--negative-prompt` and a `--guidance-scale` greater than one. CFG uses
|
||||
the ordinary linear combination without the older Qwen-Image norm correction.
|
||||
Positive and negative prompts have separate request-owned prefix caches.
|
||||
|
||||
TP uses native parallel projections. Ulysses and Ring shard target-image
|
||||
attention while keeping the condition prefix replicated. The target token
|
||||
count, `(height / 16) × (width / 16)`, must be divisible by the SP degree. Encoder
|
||||
folding shards Qwen3-VL's language projections using the native encoder TP group.
|
||||
Full-checkpoint editing passed with TP2 × Ulysses2 and TP2 × Ring2 + FlashAttention
|
||||
on four B200 GPUs. These CLI checks do not mark every HTTP topology as verified.
|
||||
|
||||
VAE tiling is disabled by default for both encoding and decoding. Enable
|
||||
`--vae-tiling true` for tiled encoding and decoding; `--vae-sp true` also distributes tiles
|
||||
across the configured GPUs. These paths use the standard VAE runtime; tiled
|
||||
decode can differ from full image decode near tile boundaries.
|
||||
|
||||
For full-image spatial parallel decode, select **Spatial shard** or pass
|
||||
`--vae-config.parallel-decode-mode spatial_shard` with at least two GPUs.
|
||||
This mode splits feature-map height, exchanges convolution halos, and gathers
|
||||
the full map for VAE attention. It does not require `--vae-tiling` or `--vae-sp`.
|
||||
Two-B200 checks cover TP2, CFG parallelism, and all-component layerwise offload.
|
||||
FP64 component comparisons match full decode; BF16 full-checkpoint output can
|
||||
differ through floating-point rounding.
|
||||
|
||||
Select **All components layerwise** or pass `--layerwise-offload-components all`
|
||||
to stream repeated blocks in the DiT, Qwen3-VL language and vision encoders, and
|
||||
VAE encoder/decoder. Full-checkpoint 512px editing passed on one B200 and on
|
||||
two B200s with TP2 plus spatial VAE decode. This setting reduces device memory
|
||||
at the cost of host-device transfers; it is not the measured default for the
|
||||
consumer-GPU recipes above.
|
||||
|
||||
Revision `f1f3366c7c` fixes CPU/GPU initialization rounding in the vision
|
||||
encoder's rotary frequencies after device transfer. On one B200, native
|
||||
1024px/40-step generation, editing, and transparent output with all-component
|
||||
layerwise offload matched resident RGBA pixels exactly. Repeated editing after
|
||||
a transparent-generation request also matched. Resident output was unchanged
|
||||
from revision `6ee35b52fb`. These checks use FlashAttention, seed 42, and CFG 1.
|
||||
|
||||
Revision `81c8c550fa` also preserves the loader's FP8 weights and FP32 rotary
|
||||
buffers when moving the whole encoder between CPU and GPU. With that fix,
|
||||
`--text-encoder-cpu-offload true` matched resident generation, editing, and
|
||||
transparent RGBA pixels for both native precision and the combined serialized
|
||||
FP8 export in the same B200 workload, including repeated editing.
|
||||
|
||||
The pipeline also supports the shared
|
||||
[disaggregated runtime](/docs/sglang-diffusion/disaggregation). The encoder role
|
||||
loads both Qwen3-VL and the VAE to prepare reference-image conditioning; nested
|
||||
condition tensors and complex RoPE tensors transfer with the request. Separate
|
||||
encoder, denoiser, and decoder processes matched monolithic RGBA output for
|
||||
512px/4-step generation, editing, different prompt lengths, and CFG on B200.
|
||||
That check used same-host Mooncake TCP; multi-host RDMA remains unverified.
|
||||
|
||||
Online FP8 is available independently for the DiT and encoder through
|
||||
`--component-quantizations.transformer fp8` and
|
||||
`--component-quantizations.text_encoder fp8`. Each component and the combination
|
||||
passed 1024px/40-step HTTP generation and editing on a resident B200. FP8 changes
|
||||
the output: in one generation/edit pair, DiT-only FP8 gave RGBA PSNR
|
||||
37.56/41.07 dB against native precision; quantizing both gave 32.66/40.99 dB.
|
||||
These samples do not establish general image or alpha quality. Native precision
|
||||
remains the default.
|
||||
|
||||
### Serialized FP8 components
|
||||
|
||||
Select a **Serialized FP8** precision option in the picker and set the component
|
||||
directories under **Variables**. The tested format is E4M3FN weights with one
|
||||
FP32 `weight_scale` per linear and dynamic activation quantization. Each
|
||||
component directory contains its own architecture `config.json`, weight shards,
|
||||
and index; merge this top-level quantization configuration into its `config.json`:
|
||||
|
||||
```json
|
||||
{
|
||||
"quantization_config": {
|
||||
"quant_method": "fp8",
|
||||
"activation_scheme": "dynamic"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Load compatible exported components through the shared loader:
|
||||
|
||||
```bash Command
|
||||
sglang serve \
|
||||
--model-path /models/qwen-image-2.1 \
|
||||
--model-id Qwen-Image-2.1 \
|
||||
--component-paths.transformer /models/qwen-image-2.1-fp8/transformer \
|
||||
--component-paths.text_encoder /models/qwen-image-2.1-fp8/text_encoder \
|
||||
--num-gpus 1 --performance-mode speed --attention-backend fa \
|
||||
--host 0.0.0.0 --port 30010
|
||||
```
|
||||
|
||||
Use either override independently, or both as shown. Omit online quantization
|
||||
flags: the component metadata selects serialized loading. Adding metadata to
|
||||
BF16 weights does not convert them. The validated export quantizes 224 DiT
|
||||
attention/MLP matrices and 252 Qwen3-VL language matrices; the vision encoder,
|
||||
embeddings, output head, other DiT projections, and VAE retain native precision.
|
||||
All 476 loaded matrices and scales matched their serialized values.
|
||||
|
||||
At revision `5a117c9f3f`, DiT-only, encoder-only, and combined exports passed
|
||||
1024px/40-step generation, editing, and transparent PNG requests on B200 with
|
||||
FlashAttention, seed 42, and CFG 1. The combined export also passed TP2 with
|
||||
encoder folding and single-GPU `--layerwise-offload-components all`.
|
||||
At that revision, offload matched resident generation and transparent output
|
||||
exactly, but editing differed at 49.50 dB RGBA PSNR. Revision `f1f3366c7c` fixes
|
||||
the vision rotary initialization difference: a new 1024px/40-step comparison
|
||||
matched resident generation, editing, and transparent RGBA pixels exactly
|
||||
with all-component layerwise offload. Resident outputs were unchanged. TP2
|
||||
still changes numerical results.
|
||||
|
||||
| Serialized FP8 scope | Generation RGBA PSNR vs native | Edit RGBA PSNR vs native |
|
||||
| --- | --- | --- |
|
||||
| DiT | 38.35 dB | 40.94 dB |
|
||||
| Encoder | 34.46 dB | 49.19 dB |
|
||||
| Both | 34.93 dB | 41.25 dB |
|
||||
|
||||
For the combined export, the transparent cat's alpha channel measured 32.03 dB
|
||||
PSNR and 0.81 mean absolute error on the 0–255 scale against native precision;
|
||||
individual boundary pixels can differ substantially. Online FP8 for both
|
||||
components also produced a real transparent PNG in this check. These are
|
||||
single-example comparisons, not a quality guarantee. Offline tensorwise scales
|
||||
differ from B200 online FP8's channelwise scales.
|
||||
|
||||
### GGUF components
|
||||
|
||||
Select **GGUF DiT**, **GGUF encoder**, or **GGUF DiT + encoder** under Server
|
||||
precision, then set the corresponding `.gguf` files under **Variables**.
|
||||
The picker uses `--component-weights-paths.transformer` and
|
||||
`--component-weights-paths.text_encoder`, retaining each component's architecture
|
||||
config from the base checkpoint. Each file must contain the entire component
|
||||
with native checkpoint tensor names. No online quantization flag is needed;
|
||||
the loader reads the quantization type from each GGUF tensor.
|
||||
|
||||
The tested Q4_0 export quantizes the same 224 DiT and 252 language-encoder
|
||||
matrices listed above. Other tensors retain native precision, including the
|
||||
vision tower, embeddings, output head, and VAE. Its DiT and encoder files are
|
||||
3.91 and 7.03 GiB respectively. All 476 loaded packed matrices matched the
|
||||
exported bytes; sampled CUDA dequantization matched the GGUF CPU reference
|
||||
after conversion to BF16.
|
||||
|
||||
At revision `7e0d4e9185`, DiT-only, encoder-only, and combined Q4_0 exports
|
||||
passed 1024px/40-step HTTP generation, editing, and transparent PNG output on
|
||||
B200 with FlashAttention, seed 42, and CFG 1. These are private validation
|
||||
exports, not published download targets. Use a compatible export of weights
|
||||
you are authorized to access.
|
||||
|
||||
The combined export also passed TP2 with encoder folding. On one GPU,
|
||||
all-component layerwise offload and whole-encoder CPU offload each matched
|
||||
resident generation, editing, and transparent RGBA pixels exactly. TP2 changed
|
||||
numerical results. Quantization itself is lossy:
|
||||
|
||||
| Q4_0 scope | Generation RGBA PSNR vs native | Edit RGBA PSNR vs native |
|
||||
| --- | --- | --- |
|
||||
| DiT | 24.99 dB | 33.66 dB |
|
||||
| Encoder | 28.97 dB | 43.26 dB |
|
||||
| Both | 23.86 dB | 33.46 dB |
|
||||
|
||||
The combined export's transparent cat retained alpha values from 0 to 255,
|
||||
with 66.8% of pixels at alpha 5 or below. Against native precision, its alpha
|
||||
PSNR was 21.20 dB and mean absolute error was 3.29/255; individual boundary
|
||||
pixels differed by up to 255. These single-example comparisons do not establish
|
||||
general image or cutout quality. Keep native precision when exact output is
|
||||
required.
|
||||
|
||||
GGUF reduces weight storage; it is not a promise of lower latency. The runtime
|
||||
dequantizes packed linears before BF16 matrix multiplication. Other GGUF tensor
|
||||
types, exports, and hardware need separate validation.
|
||||
See the shared [GGUF guide](/docs/sglang-diffusion/quantization#gguf)
|
||||
for loader and parallelism constraints.
|
||||
|
||||
### NVFP4 components
|
||||
|
||||
Select **NVFP4 DiT**, **NVFP4 encoder**, or **NVFP4 DiT + encoder** in the
|
||||
picker, then set the component directories under **Variables**. These options
|
||||
require Blackwell; H200 and RTX 4090 cannot run this native FP4 path. B200 has
|
||||
completed the checks below. RTX PRO 6000 and RTX 5090 remain unverified for this
|
||||
model's NVFP4 exports; their FlashInfer backend defaults to `auto`, because
|
||||
TensorRT-LLM FP4 GEMM does not support SM120. Keep that default on these GPUs.
|
||||
|
||||
Each exported directory contains its architecture config, weight shards, and
|
||||
index. The config declares `quant_method: modelopt`, `quant_algo: NVFP4`, and
|
||||
block size 16, with exclusions for native-precision layers. Use
|
||||
`--component-paths.transformer` and/or `--component-paths.text_encoder` to load
|
||||
the exported directories. Omit online quantization flags; metadata alone does
|
||||
not convert native weights into an NVFP4 checkpoint.
|
||||
|
||||
The private validation export quantizes the same 224 DiT and 252 language
|
||||
matrices as the FP8 example. Vision, embeddings, the output head, other DiT
|
||||
projections, and VAE retain native precision. Weight quantization uses ModelOpt
|
||||
0.46.1 with max calibration; static activation scales come from six separate
|
||||
1024px/40-step requests, including two edits and one transparent generation.
|
||||
This small calibration set does not establish general quality. It does not
|
||||
use SVDQuant or AWQ. All 476 loaded packed weights, block scales, and global
|
||||
scales matched the export after the runtime's layout transforms.
|
||||
|
||||
At revision `57b625d3e3`, each component and both together passed 1024px/40-step
|
||||
HTTP generation, editing, and transparent PNG output on B200 with
|
||||
FlashAttention, seed 42, CFG 1, and FlashInfer TensorRT-LLM FP4 GEMM. The combined
|
||||
export also passed TP2 with encoder folding. Single-GPU all-component layerwise
|
||||
offload and whole-encoder CPU offload each matched the combined resident RGBA
|
||||
pixels exactly. TP2 changed numerical results.
|
||||
|
||||
| NVFP4 scope | Generation RGBA PSNR vs native | Edit RGBA PSNR vs native |
|
||||
| --- | --- | --- |
|
||||
| DiT | 24.97 dB | 31.56 dB |
|
||||
| Encoder | 26.48 dB | 36.63 dB |
|
||||
| Both | 19.36 dB | 29.96 dB |
|
||||
|
||||
The combined export's transparent cat retained alpha from 0 to 255, with
|
||||
67.8% of pixels at alpha 5 or below. Against native precision, alpha PSNR was
|
||||
23.81 dB and mean absolute error was 2.22/255; some boundary pixels differed
|
||||
by 255. These are single-example comparisons of private exports, not download
|
||||
targets or quality guarantees. Native precision remains the default. See the
|
||||
shared [NVFP4 guide](/docs/sglang-diffusion/quantization#modelopt-nvfp4) for loader
|
||||
details.
|
||||
|
||||
### LoRA and execution options
|
||||
|
||||
LoRA uses the shared `--lora-path` and `--lora-merge-mode dynamic|merge` options
|
||||
and runtime adapter APIs. Diffusers keys prefixed with `transformer.` map to
|
||||
the native DiT. A synthetic adapter covering attention and MLP projections
|
||||
passed dynamic loading, merging, and removal on one B200 and TP2 with encoder
|
||||
folding. Both removal paths restored the base image exactly. This verifies
|
||||
adapter application and lifecycle, not the quality of a trained LoRA.
|
||||
|
||||
Cache-DiT hooks operate on target-image transformer blocks. Breakable CUDA
|
||||
Graph execution fills each request's prefix caches eagerly, then replays
|
||||
matching warmup graphs with those cache tensors as inputs. Warmup and request
|
||||
condition-prefix lengths must match, in addition to the output resolution;
|
||||
unseen shapes run eagerly. Text buckets alone cannot pad condition KV without
|
||||
changing attention semantics. FlashAttention, Sage
|
||||
attention and Torch SDPA are wired through the native attention layers;
|
||||
causal text runs use exact masked SDPA. Sage and Cache-DiT can change numerical
|
||||
results and require application-specific quality checks.
|
||||
|
||||
See the [compatibility inventory](/docs/sglang-diffusion/compatibility_matrix)
|
||||
for tested configurations and remaining validation boundaries. These checks
|
||||
are functional and numerical comparisons. The platform measurements above cover
|
||||
their stated HTTP workload; broader image quality is not evaluated.
|
||||
@@ -28,7 +28,7 @@ Image models generate one image request as a bounded denoising job, usually with
|
||||
<Card
|
||||
title="Qwen-Image"
|
||||
mode="card"
|
||||
href="/cookbook/diffusion/Qwen-Image/Qwen-Image"
|
||||
href="/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1"
|
||||
img="/cards/logos/qwen.png"
|
||||
/>
|
||||
<Card
|
||||
|
||||
@@ -1509,6 +1509,7 @@
|
||||
{
|
||||
"group": "Qwen-Image",
|
||||
"pages": [
|
||||
"cookbook/diffusion/Qwen-Image/Qwen-Image-2.1",
|
||||
"cookbook/diffusion/Qwen-Image/Qwen-Image",
|
||||
"cookbook/diffusion/Qwen-Image/Qwen-Image-Edit"
|
||||
]
|
||||
|
||||
@@ -5,6 +5,122 @@ description: "Browse model families and public checkpoints supported by SGLang D
|
||||
|
||||
import { DiffusionModelCatalog } from '/src/snippets/diffusion/model-catalog.jsx';
|
||||
|
||||
## Qwen-Image 2.1 integration status
|
||||
|
||||
Qwen-Image 2.1 has a separate native pipeline for text-to-image and optional
|
||||
reference-image conditioning. Full-checkpoint validation on H200 includes
|
||||
1024×1024, 40-step generation and editing. Reference comparisons use
|
||||
Transformers 4.57.3, BF16 and matching Torch SDPA settings with cuDNN SDPA
|
||||
disabled. The native encoder explicitly preserves that version's pre-final-norm
|
||||
hidden state and vision position interpolation semantics.
|
||||
|
||||
Full-checkpoint functional checks on up to two H200 GPUs cover DiT TP, Ulysses,
|
||||
Ring + FA, CFG parallel, DiT layerwise offload, FA/Sage, parallel
|
||||
tiled VAE decode, Cache-DiT, breakable CUDA graphs, and multiple reference
|
||||
images and outputs. One H200 with resident weights also passed HTTP generation
|
||||
and editing, including multiple outputs and reference images. These are
|
||||
functional checks, not a broad quality evaluation of lossy optimizations.
|
||||
|
||||
Additional full-checkpoint checks on B200 cover native encoder TP, combined
|
||||
TP2 × Ulysses2 and TP2 × Ring2 + FA on four GPUs, and
|
||||
`--layerwise-offload-components all`. Two-GPU VAE `spatial_shard` decode passed
|
||||
with TP2, CFG parallelism, and all-component offload. These use 512px/4-step
|
||||
editing, plus a 1024px/4-step comparison of spatial and full VAE decode.
|
||||
Single-GPU 1024px/40-step generation, editing, and transparent PNG output also
|
||||
passed at source revision `d629d234f1`.
|
||||
|
||||
Two-B200 component regressions assert actual encoder weight sharding and compare
|
||||
text and image conditioning with a single-rank reference. VAE regressions cover
|
||||
both decoder structures and even/odd latent heights in FP32 and FP64. FP64
|
||||
spatial decode matches the full decoder; BF16 full-checkpoint output is not
|
||||
bit-exact. In the tested image, spatial versus full decode gave RGBA PSNR
|
||||
60.69 dB; this is one sample, not a general quality guarantee.
|
||||
|
||||
Request-local prefix KV reuse and new-request graph replay have CUDA regression
|
||||
tests. The deployment picker marks only its exact tested HTTP combinations as
|
||||
verified, including H200, B200, RTX PRO 6000 96GB, RTX 5090, and RTX 4090.
|
||||
CLI-only combinations remain Unverified in the picker.
|
||||
|
||||
On `1eab5de5990`, one RTX PRO 6000 Blackwell Server Edition passed repeated
|
||||
1024px/40-step HTTP generation, editing, transparent generation, and transparent
|
||||
input editing with resident weights or DiT layerwise offload and Torch SDPA.
|
||||
Repeated requests and corresponding outputs across placements matched RGBA
|
||||
pixels; transparent outputs retained alpha from 0 to 255. Its SM120
|
||||
RoPE, normalization, fusion-gate, and model regressions also passed. This check
|
||||
does not establish multi-GPU or quantized-checkpoint coverage on that platform.
|
||||
|
||||
Disaggregated encoder, denoiser, and decoder roles passed full-checkpoint HTTP
|
||||
generation, variable-length prompts, image editing, and CFG on three B200 GPUs,
|
||||
using Mooncake's same-host TCP transport. At 512px/4 steps, their RGBA pixels
|
||||
matched the monolithic server exactly. Multi-host RDMA and multi-rank roles
|
||||
are not covered by that check.
|
||||
|
||||
Online FP8 for the DiT, encoder, and both together passed 1024px/40-step HTTP
|
||||
generation and editing on one B200; both together also passed transparent PNG
|
||||
generation. Serialized E4M3FN FP8 component folders with scalar weight scales
|
||||
and dynamic activations passed generation, editing, and transparent output for
|
||||
each component and both together. The combined export also passed separate
|
||||
checks with TP2 plus encoder folding and with single-GPU all-component
|
||||
layerwise offload. The 476 loaded quantized matrices and
|
||||
scales matched their checkpoint values. Quantization and TP change numerical
|
||||
results. Revision `f1f3366c7c` fixes vision rotary initialization rounding during
|
||||
CPU/GPU transfer. In 1024px/40-step B200 checks, all-component layerwise offload
|
||||
matched resident generation, editing, and transparent RGBA pixels exactly
|
||||
for both native precision and the combined serialized export. Revision
|
||||
`81c8c550fa` additionally preserves FP8 weights and FP32 buffers during
|
||||
whole-encoder CPU offload; separate checks of that mode also matched resident
|
||||
pixels. Resident output was unchanged by these fixes. See the cookbook for
|
||||
sample quality metrics.
|
||||
|
||||
Native-name Q4_0 GGUF files passed 1024px/40-step B200 HTTP generation,
|
||||
editing, and transparent output for DiT-only, encoder-only, and both together.
|
||||
The 476 packed matrices matched their exported bytes. The combined export also
|
||||
passed TP2 with encoder folding; separate all-component layerwise and
|
||||
whole-encoder CPU offload checks matched resident RGBA pixels exactly.
|
||||
Quantization and TP change output pixels, including alpha. These checks cover
|
||||
the documented private exports, not arbitrary community checkpoints. Other
|
||||
GGUF types and hardware remain unverified for this model.
|
||||
|
||||
Calibrated ModelOpt-format NVFP4 component directories also passed the same
|
||||
six-configuration B200 HTTP matrix: each component, both together, two offload
|
||||
modes, and TP2 with encoder folding. All 476 packed matrices and their scales
|
||||
matched the exports after runtime layout transforms. Both offload modes
|
||||
matched resident RGBA pixels exactly. NVFP4 and TP2 change image and alpha
|
||||
values; this private export uses max calibration on only six requests, without
|
||||
SVDQuant or AWQ. See the cookbook for measured errors. Native FP4 requires
|
||||
Blackwell, and this model's NVFP4 exports remain unverified on RTX PRO 6000 and
|
||||
RTX 5090.
|
||||
|
||||
Synthetic Diffusers-format LoRA passed dynamic loading, merging,
|
||||
and removal on one B200 and TP2 with encoder folding; both removal paths
|
||||
restored the base image exactly. Trained-adapter quality remains unverified.
|
||||
Platform latency measurements
|
||||
and their exact workloads are in the cookbook. Breakable graphs
|
||||
replay only matching warmup signatures, including condition-prefix length;
|
||||
other shapes run eagerly. VAE tiling is opt-in and can change numerical results.
|
||||
Do not infer compatibility from the older Qwen-Image row.
|
||||
|
||||
Use an authorized local checkpoint with `--model-id Qwen-Image-2.1`.
|
||||
See the [Qwen-Image 2.1 cookbook](/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1)
|
||||
for checkpoint layout and usage. This entry does not assert public weight
|
||||
availability. The standard two-GPU E2E suite includes `qwen_image21_t2i_tp2`
|
||||
with TP2, 1024px/40-step generation, two requests, and image consistency checks.
|
||||
Its runners need access to `Qwen/Qwen-Image-2.1`.
|
||||
|
||||
The additional opt-in HTTP cases check repeated generation, editing, and real
|
||||
RGBA alpha output from a local checkpoint:
|
||||
|
||||
```bash
|
||||
export SGLANG_QWEN_IMAGE21_TEST_MODEL=/models/Qwen-Image-2.1
|
||||
export SGLANG_QWEN_IMAGE21_TEST_IMAGE=/path/to/reference.png
|
||||
python -m pytest -q python/sglang/multimodal_gen/test/server/test_server_qwen_image21.py
|
||||
```
|
||||
|
||||
The module skips without the model variable; only editing skips when the image
|
||||
variable is absent. Credentials stay in the runner environment.
|
||||
|
||||
## Supported model inventory
|
||||
|
||||
Use a listed checkpoint as `--model-path` with `sglang generate` or
|
||||
`sglang serve`. This registry-backed list contains known public entry points;
|
||||
family detection may also support compatible local directories. Open the linked
|
||||
|
||||
@@ -774,6 +774,14 @@ sglang generate \
|
||||
--save-output
|
||||
```
|
||||
|
||||
Qwen-Image 2.1 also supports independent NVFP4 DiT and native language-encoder
|
||||
component directories through `--component-paths.transformer` and
|
||||
`--component-paths.text_encoder`. Private max-calibrated exports passed B200
|
||||
generation, editing, transparent RGBA, offload, and TP2 checks. These checks
|
||||
do not cover arbitrary exports or RTX 5090. See its
|
||||
[cookbook](/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1#nvfp4-components) for
|
||||
calibration scope and image/alpha error measurements.
|
||||
|
||||
### Notes
|
||||
|
||||
- Use `--transformer-path` for mixed ModelOpt NVFP4 transformer repos or local
|
||||
@@ -784,8 +792,9 @@ sglang generate \
|
||||
primary `--transformer-path` override targets only `transformer`. Use a
|
||||
per-component override such as `--transformer-2-path` only when you
|
||||
intentionally want a non-default `transformer_2`.
|
||||
- On Blackwell, the diffusion ModelOpt NVFP4 path defaults to FlashInfer
|
||||
TensorRT-LLM FP4 GEMM (`flashinfer_trtllm`).
|
||||
- On B200, the diffusion ModelOpt NVFP4 path defaults to FlashInfer
|
||||
TensorRT-LLM FP4 GEMM (`flashinfer_trtllm`). On SM120, including RTX 5090,
|
||||
the default is `auto`; TensorRT-LLM FP4 GEMM does not support that architecture.
|
||||
- The published Qwen Image NVFP4 exports keep the `img_mod`/`txt_mod`
|
||||
modulation projections and first/last transformer blocks in BF16.
|
||||
- Qwen-Image NVFP4 does not always improve latency at 1024x1024. On B200, the
|
||||
@@ -916,13 +925,21 @@ startup rejects a degree that cuts a row-parallel matrix inside a GGML block.
|
||||
| MiniMax-H3 `fl2va`, pruned AdaLN curve | [`unsloth/MiniMax-H3-GGUF`](https://huggingface.co/unsloth/MiniMax-H3-GGUF) `minimax_h3_fl2va_pruned-Q4_K.gguf` (10.7 GiB loaded DiT) | 1x H200 (141 GiB) | 2-step t2va 1344x768, 107 frames, H.264 + AAC |
|
||||
| MiniMax-H3 `fl2va`, pruned AdaLN curve | `minimax_h3_fl2va_pruned-Q4_K.gguf` | 1x GB300 (CUDA 13, PyTorch 2.13) | 50-step t2va 1344x768, 107 frames, H.264 + AAC; 105.38 s and 80.88 GB peak |
|
||||
| MiniMax-H3 `fl2va`, pruned AdaLN curve | `minimax_h3_fl2va_pruned-Q4_K.gguf` | 2x GB300, TP2 (CUDA 13, PyTorch 2.13) | 2-step t2va 1344x768, 107 frames, H.264 + AAC; 7.55 s and 51.90 GB peak per rank |
|
||||
| Qwen-Image 2.1 | Private native-name Q4_0 exports; DiT 3.91 GiB, encoder 7.03 GiB | 1x B200; separate TP2 check | 1024px/40-step generation, editing, and transparent RGBA; each component and both together; combined offload matched resident pixels |
|
||||
|
||||
The DiT loads at 17.5 GiB against 61.7 GiB for the BF16 checkpoint. Weight
|
||||
The unpruned MiniMax-H3 DiT loads at 17.5 GiB against 61.7 GiB for the BF16 checkpoint. Weight
|
||||
fidelity was checked tensor-by-tensor against the BF16 reference: cosine
|
||||
1.00000 for the F32/BF16 tensors and 0.9973 for Q4_K/Q4_0.
|
||||
|
||||
Not validated in the measurements above: any other quantization type, the
|
||||
`ref2va` partition, and a BF16-vs-GGUF output comparison.
|
||||
The MiniMax-H3 measurements do not validate other quantization types, the
|
||||
`ref2va` partition, or a BF16-vs-GGUF output comparison.
|
||||
|
||||
Qwen-Image 2.1's Q4_0 exports use the base component configs and native tensor
|
||||
names. Its [cookbook](/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1#gguf-components)
|
||||
documents the tested matrix selection, offload modes, TP2, and image/alpha
|
||||
differences against native precision. These private exports establish loading
|
||||
and execution compatibility, not general output quality or compatibility with
|
||||
other GGUF exports.
|
||||
|
||||
## Nunchaku (SVDQuant)
|
||||
|
||||
|
||||
@@ -1277,6 +1277,7 @@ export const Deployment = ({ config, benchmarks }) => {
|
||||
};
|
||||
|
||||
const [sel, setSel] = useState(() => initialSelectionFromCells());
|
||||
const [selectionHydrated, setSelectionHydrated] = useState(false);
|
||||
const INTERNAL_HASH_STATE_KEY = "__sglangDeployInternalHash";
|
||||
const DEPLOYMENT_COMPONENT_ID = "deployment-configurator";
|
||||
useEffect(() => {
|
||||
@@ -1310,12 +1311,14 @@ export const Deployment = ({ config, benchmarks }) => {
|
||||
if (el) el.scrollIntoView({ behavior: "smooth", block: "start" });
|
||||
};
|
||||
hydrate();
|
||||
setSelectionHydrated(true);
|
||||
window.addEventListener("hashchange", hydrate);
|
||||
return () => window.removeEventListener("hashchange", hydrate);
|
||||
}, []);
|
||||
// history.replaceState does NOT fire hashchange — dispatch a custom event so
|
||||
// the Playground hears chip-click selection changes.
|
||||
useEffect(() => {
|
||||
if (!selectionHydrated) return;
|
||||
const target = "#" + new URLSearchParams(sel).toString();
|
||||
if (window.location.hash !== target) {
|
||||
const historyState =
|
||||
@@ -1329,7 +1332,7 @@ export const Deployment = ({ config, benchmarks }) => {
|
||||
);
|
||||
}
|
||||
window.dispatchEvent(new CustomEvent("sglang-deploy-sel", { detail: sel }));
|
||||
}, [sel]);
|
||||
}, [sel, selectionHydrated]);
|
||||
|
||||
const [modal, setModal] = useState(null); // 'curl' | 'env' | 'bench' | null
|
||||
useEffect(() => {
|
||||
|
||||
@@ -0,0 +1,386 @@
|
||||
export const config = (() => {
|
||||
const sm120Hardware = ["rtx5090", "rtxpro6000"];
|
||||
const platformAttention = (s) => sm120Hardware.includes(s.hw) ? "sdpa" : "fa";
|
||||
const effectiveAttention = (s) => s.attention === "platform" || (sm120Hardware.includes(s.hw) && s.attention === "fa") ? platformAttention(s) : s.attention;
|
||||
|
||||
const config = {
|
||||
modelName: "Qwen-Image 2.1",
|
||||
supportedHardware: ["h200", "b200", "rtxpro6000", "rtx5090", "rtx4090"],
|
||||
hardware: [
|
||||
{ id: "rtxpro6000", label: "RTX PRO 6000", vram: "96GB", vendor: "consumer" },
|
||||
{ id: "rtx5090", label: "RTX 5090", vram: "32GB", vendor: "consumer" },
|
||||
{ id: "rtx4090", label: "RTX 4090", vram: "24GB", vendor: "consumer" },
|
||||
],
|
||||
groupHardware: false,
|
||||
matchDims: [],
|
||||
|
||||
overlayDims: [
|
||||
{
|
||||
id: "weights",
|
||||
title: "Checkpoint weights",
|
||||
scope: "base",
|
||||
description: "One checkpoint serves generation and editing. Set its authorized local path under Variables.",
|
||||
default: "default",
|
||||
options: [{ id: "default", label: "Qwen-Image 2.1", flags: [] }],
|
||||
},
|
||||
{
|
||||
id: "mode",
|
||||
title: "Request mode",
|
||||
scope: "base",
|
||||
description: "Switch between JSON generation and PNG uploads to the image-edit endpoint.",
|
||||
default: "text",
|
||||
options: [
|
||||
{ id: "text", label: "Text to image" },
|
||||
{ id: "edit", label: "Image edit", description: "Upload one reference PNG, preserving its alpha channel." },
|
||||
{ id: "multi", label: "Multi-image edit", description: "Upload two ordered references; Picture 1 and Picture 2 follow this order." },
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "placement",
|
||||
title: "Placement",
|
||||
scope: "serve",
|
||||
description: "Hardware selection applies its recommended placement. Stream DiT layers when the full pipeline exceeds device memory.",
|
||||
learnMore: "#5-runtime-features",
|
||||
default: "resident",
|
||||
options: [
|
||||
{
|
||||
id: "resident", label: "Resident",
|
||||
recommendedWhen: (s) => ["h200", "b200", "rtxpro6000"].includes(s.hw),
|
||||
disabled: (s) => ["rtx5090", "rtx4090"].includes(s.hw) && Number(s.gpus_per_node) === 1,
|
||||
disableReason: "The full resident pipeline exceeds one consumer GPU's memory. Select CPU offload.",
|
||||
flags: (s) => [Number(s.gpus_per_node) === 1 ? "--performance-mode speed" : "--performance-mode manual"],
|
||||
description: "Keep all components on the GPU. Recommended for H200, B200, and RTX PRO 6000 96GB. RTX 5090 and RTX 4090 need offload.",
|
||||
},
|
||||
{
|
||||
id: "offload", label: "CPU offload",
|
||||
flags: (s) => ["--performance-mode manual", "--dit-layerwise-offload true", ...(s.hw === "rtx4090" ? ["--text-encoder-cpu-offload true"] : [])],
|
||||
recommendedWhen: (s) => ["rtx5090", "rtx4090"].includes(s.hw),
|
||||
soft: (s) => !["rtxpro6000", "rtx5090", "rtx4090"].includes(s.hw) || Number(s.gpus_per_node) !== 1,
|
||||
softReason: "This offload topology has not completed an HTTP verification run.",
|
||||
description: "Streams DiT layers. RTX 4090 also offloads the encoder between requests to leave room for image editing. Requires sufficient host RAM.",
|
||||
},
|
||||
{
|
||||
id: "all_offload", label: "All components layerwise",
|
||||
flags: ["--performance-mode manual", "--layerwise-offload-components all"],
|
||||
soft: true, softReason: "Full-checkpoint 512px editing passed on B200, including TP2 with spatial VAE decode; this HTTP recipe is unverified.",
|
||||
description: "Streams repeated blocks in the DiT, Qwen3-VL, and VAE. Uses more host-device transfers to reduce device memory.",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "attention",
|
||||
title: "Attention",
|
||||
scope: "serve",
|
||||
description: "Choose the target-image attention kernel. Text attention retains its causal mask.",
|
||||
learnMore: "#5-runtime-features",
|
||||
default: "platform",
|
||||
options: [
|
||||
{
|
||||
id: "platform", label: "Automatic", recommended: true,
|
||||
flags: (s) => [`--attention-backend ${platformAttention(s) === "sdpa" ? "torch_sdpa" : "fa"}`],
|
||||
description: "Uses SDPA on RTX PRO 6000 and RTX 5090, and FlashAttention on the other listed GPUs.",
|
||||
},
|
||||
{ id: "fa", label: "FlashAttention", flags: ["--attention-backend fa"], description: "Exact attention with a fused kernel. This runtime falls back to Torch SDPA on RTX PRO 6000 and RTX 5090." },
|
||||
{
|
||||
id: "sdpa", label: "Torch SDPA", flags: ["--attention-backend torch_sdpa"],
|
||||
soft: (s) => !config.commandBuilder.resource.verifiedRecipes.some((r) => r.hw === s.hw && r.placement === s.placement && r.attentions.includes("sdpa") && Number(s.gpus_per_node) === r.gpus_per_node),
|
||||
softReason: "This hardware and placement combination has not completed HTTP verification with SDPA.",
|
||||
description: "Use for reference comparisons. Floating-point reduction order can differ from FlashAttention.",
|
||||
},
|
||||
{
|
||||
id: "sage", label: "SageAttention", flags: ["--attention-backend sage_attn"],
|
||||
soft: true, softReason: "CLI smoke test passed; image and alpha quality need workload-specific validation.",
|
||||
description: "Approximate attention; requires the SageAttention dependency.",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "precision",
|
||||
title: "Precision",
|
||||
scope: "serve",
|
||||
description: "Native precision is the default. Quantization changes image and alpha values. Set compatible FP8/NVFP4 directories or GGUF files under Variables.",
|
||||
default: "native",
|
||||
options: [
|
||||
{ id: "native", label: "Native BF16 / FP32", recommended: true },
|
||||
{
|
||||
id: "fp8_dit", label: "Online FP8 DiT", flags: ["--component-quantizations.transformer fp8"],
|
||||
soft: true, softReason: "Online FP8 passed 1024px/40-step generation and editing on one resident B200. Other hardware, alpha, and feature combinations remain unverified.",
|
||||
},
|
||||
{
|
||||
id: "fp8_encoder", label: "Online FP8 encoder", flags: ["--component-quantizations.text_encoder fp8"],
|
||||
soft: true, softReason: "Online encoder FP8 passed 1024px/40-step generation and editing on one resident B200. It changes conditioning and output pixels.",
|
||||
},
|
||||
{
|
||||
id: "fp8_both", label: "Online FP8 DiT + encoder", flags: ["--component-quantizations.transformer fp8", "--component-quantizations.text_encoder fp8"],
|
||||
soft: true, softReason: "Online FP8 for both components passed generation, editing, and transparent output on one resident B200. Quality depends on the workload.",
|
||||
},
|
||||
{
|
||||
id: "serialized_fp8_dit", label: "Serialized FP8 DiT", flags: ['--component-paths.transformer "{{FP8_DIT_PATH}}"'],
|
||||
soft: true, softReason: "A tensorwise E4M3FN component export passed 1024px/40-step generation, editing, and transparent output on B200. Validate your exported checkpoint's quality.",
|
||||
},
|
||||
{
|
||||
id: "serialized_fp8_encoder", label: "Serialized FP8 encoder", flags: ['--component-paths.text_encoder "{{FP8_ENCODER_PATH}}"'],
|
||||
soft: true, softReason: "A tensorwise E4M3FN language encoder export passed generation, editing, and transparent output on B200; vision weights retain native precision.",
|
||||
},
|
||||
{
|
||||
id: "serialized_fp8_both", label: "Serialized FP8 DiT + encoder", flags: ['--component-paths.transformer "{{FP8_DIT_PATH}}"', '--component-paths.text_encoder "{{FP8_ENCODER_PATH}}"'],
|
||||
soft: true, softReason: "Exported components passed 1024px/40-step generation, editing, and transparent output on B200. All-component offload matched resident pixels after the vision RoPE fix; TP2 changes numerical results. Validate your exported checkpoint's quality.",
|
||||
},
|
||||
{
|
||||
id: "gguf_dit", label: "GGUF DiT", flags: ['--component-weights-paths.transformer "{{GGUF_DIT_PATH}}"'],
|
||||
soft: true, softReason: "A Q4_0 DiT export passed 1024px/40-step generation, editing, and transparent output on B200. Other exports and hardware need validation.",
|
||||
},
|
||||
{
|
||||
id: "gguf_encoder", label: "GGUF encoder", flags: ['--component-weights-paths.text_encoder "{{GGUF_ENCODER_PATH}}"'],
|
||||
soft: true, softReason: "A native-name Q4_0 language encoder export passed generation, editing, and transparent output on B200; vision weights retain native precision.",
|
||||
},
|
||||
{
|
||||
id: "gguf_both", label: "GGUF DiT + encoder", flags: ['--component-weights-paths.transformer "{{GGUF_DIT_PATH}}"', '--component-weights-paths.text_encoder "{{GGUF_ENCODER_PATH}}"'],
|
||||
soft: true, softReason: "Combined Q4_0 exports passed 1024px/40-step generation, editing, and transparent output on B200. GGUF reduces weight memory; output quality and speed depend on the export and workload.",
|
||||
},
|
||||
{
|
||||
id: "nvfp4_dit", label: "NVFP4 DiT", flags: ['--component-paths.transformer "{{NVFP4_DIT_PATH}}"'],
|
||||
disabled: (s) => !["b200", "rtxpro6000", "rtx5090"].includes(s.hw),
|
||||
disableReason: "Native NVFP4 requires a Blackwell GPU (compute capability 10.0 or newer).",
|
||||
soft: true, softReason: "A calibrated ModelOpt-format DiT export passed 1024px/40-step generation, editing, and transparent output on B200. Other exports, RTX PRO 6000, and RTX 5090 need validation.",
|
||||
},
|
||||
{
|
||||
id: "nvfp4_encoder", label: "NVFP4 encoder", flags: ['--component-paths.text_encoder "{{NVFP4_ENCODER_PATH}}"'],
|
||||
disabled: (s) => !["b200", "rtxpro6000", "rtx5090"].includes(s.hw),
|
||||
disableReason: "Native NVFP4 requires a Blackwell GPU (compute capability 10.0 or newer).",
|
||||
soft: true, softReason: "A calibrated language-encoder export passed generation, editing, and transparent output on B200; vision weights retain native precision. Output quality requires validation.",
|
||||
},
|
||||
{
|
||||
id: "nvfp4_both", label: "NVFP4 DiT + encoder", flags: ['--component-paths.transformer "{{NVFP4_DIT_PATH}}"', '--component-paths.text_encoder "{{NVFP4_ENCODER_PATH}}"'],
|
||||
disabled: (s) => !["b200", "rtxpro6000", "rtx5090"].includes(s.hw),
|
||||
disableReason: "Native NVFP4 requires a Blackwell GPU (compute capability 10.0 or newer).",
|
||||
soft: true, softReason: "Combined exports passed generation, editing, transparent output, offload, and TP2 on B200. The small max-calibration sample changes image and alpha values; validate your exported checkpoint.",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "encoder",
|
||||
title: "Encoder",
|
||||
scope: "serve",
|
||||
description: "Schedule Qwen3-VL independently of target-image attention.",
|
||||
learnMore: "#5-runtime-features",
|
||||
default: "auto",
|
||||
options: [
|
||||
{ id: "auto", label: "Auto", flags: ["--encoder-parallel auto"], recommended: true },
|
||||
{ id: "replicate", label: "Replicate", flags: ["--encoder-parallel replicate"], soft: true, softReason: "Explicit replication has not been verified for this server recipe." },
|
||||
{ id: "fold", label: "Fold", flags: ["--encoder-parallel fold"], soft: true, softReason: "Native encoder TP and full-checkpoint TP2 × SP2 editing passed on B200. Requires node-local P2P; this HTTP recipe is unverified." },
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "vae",
|
||||
title: "VAE decoding",
|
||||
scope: "serve",
|
||||
description: "Decode RGBA in full, in tiles, or with spatial work distributed across GPUs.",
|
||||
learnMore: "#5-runtime-features",
|
||||
default: "full",
|
||||
options: [
|
||||
{ id: "full", label: "Full image", recommended: true, description: "Default for generation and condition-image encoding." },
|
||||
{ id: "tiled", label: "Tiled", flags: ["--vae-tiling true"], soft: true, softReason: "Repeated 512px HTTP edits passed; other tiled workloads remain unverified.", description: "Reduces activation memory; can change pixels near tile boundaries." },
|
||||
{
|
||||
id: "parallel", label: "Parallel tiles", flags: ["--vae-tiling true", "--vae-sp true"],
|
||||
disabled: (s) => Number(s.gpus_per_node) < 2,
|
||||
disableReason: "Select two GPUs before distributing VAE tiles.",
|
||||
soft: true, softReason: "Two-H200 CLI decoding passed; this HTTP recipe is unverified.",
|
||||
},
|
||||
{
|
||||
id: "spatial", label: "Spatial shard", flags: ["--vae-config.parallel-decode-mode spatial_shard"],
|
||||
disabled: (s) => Number(s.gpus_per_node) < 2,
|
||||
disableReason: "Select at least two GPUs for spatial VAE decode.",
|
||||
soft: true, softReason: "Two-B200 full-checkpoint decoding passed with TP, CFG parallelism, and all-component offload; this HTTP recipe is unverified.",
|
||||
description: "Splits feature-map height and exchanges convolution halos. Preserves full-image attention; floating-point rounding can change pixels.",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "execution",
|
||||
title: "Execution",
|
||||
scope: "serve",
|
||||
description: "Graph replay requires matching resolution and condition-prefix length.",
|
||||
learnMore: "#5-runtime-features",
|
||||
default: "eager",
|
||||
options: [
|
||||
{ id: "eager", label: "Eager", recommended: true },
|
||||
{
|
||||
id: "bcg", label: "Breakable CUDA Graph",
|
||||
flags: ["--enable-breakable-cuda-graph true", "--warmup-resolutions 512x512", "--bcg-text-buckets 64"],
|
||||
soft: true, softReason: "Only a matching 512px CLI warmup was verified. Other prompts or image prefixes can fall back to eager.",
|
||||
description: "Captures a 512px warmup. Text buckets do not pad condition KV; this is not a guaranteed replay recipe.",
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "background",
|
||||
title: "Background",
|
||||
scope: "request",
|
||||
description: "Both choices save PNG. Transparency is requested in the prompt, not imposed by postprocessing.",
|
||||
learnMore: "#transparent-png-output",
|
||||
default: "scene",
|
||||
options: [
|
||||
{ id: "scene", label: "Scene", recommended: true },
|
||||
{ id: "transparent", label: "Transparent / alpha", description: "Generate an isolated subject, or preserve the reference image's transparent background." },
|
||||
],
|
||||
},
|
||||
{
|
||||
id: "resolution",
|
||||
title: "Resolution",
|
||||
scope: "request",
|
||||
description: "Square output canvas; reference images keep their own aspect ratios.",
|
||||
default: "1024",
|
||||
options: [{ id: "512", label: "512 × 512" }, { id: "1024", label: "1024 × 1024", recommended: true }],
|
||||
},
|
||||
{
|
||||
id: "steps",
|
||||
title: "Denoising steps",
|
||||
scope: "request",
|
||||
description: "40 is the checkpoint default. Fewer steps trade detail for latency.",
|
||||
kind: "number", min: 1, max: 100, unit: "steps", default: 40, options: [],
|
||||
},
|
||||
{
|
||||
id: "outputs",
|
||||
title: "Outputs",
|
||||
scope: "request",
|
||||
description: "Generate independent images for the same prompt.",
|
||||
kind: "number", min: 1, max: 10, unit: "outputs per prompt", default: 1, options: [],
|
||||
},
|
||||
],
|
||||
|
||||
commandBuilder: {
|
||||
defaultSelection: {
|
||||
hw: "h200", nodes: 1, gpus_per_node: 1, topology_mode: "auto",
|
||||
tp_size: 1, ulysses_degree: 1, ring_degree: 1,
|
||||
},
|
||||
resource: {
|
||||
limits: { nodes: { min: 1, max: 1 }, gpus_per_node: { min: 1, max: 4 } },
|
||||
verifiedRecipes: [
|
||||
{ id: "h200-1-resident", hw: "h200", nodes: 1, gpus_per_node: 1, placement: "resident", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["fa"], default: true },
|
||||
{ id: "b200-1-resident", hw: "b200", nodes: 1, gpus_per_node: 1, placement: "resident", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["fa", "sdpa"], default: true },
|
||||
{ id: "rtxpro6000-1-resident", hw: "rtxpro6000", nodes: 1, gpus_per_node: 1, placement: "resident", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["sdpa"], default: true },
|
||||
{ id: "rtxpro6000-1-offload", hw: "rtxpro6000", nodes: 1, gpus_per_node: 1, placement: "offload", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["sdpa"] },
|
||||
{ id: "rtx5090-1-offload", hw: "rtx5090", nodes: 1, gpus_per_node: 1, placement: "offload", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["sdpa"], default: true },
|
||||
{ id: "rtx4090-1-offload", hw: "rtx4090", nodes: 1, gpus_per_node: 1, placement: "offload", tp_size: 1, ulysses_degree: 1, ring_degree: 1, encoder: "auto", attentions: ["fa"], default: true },
|
||||
],
|
||||
autoTopology: (s) => ({ tp_size: 1, ulysses_degree: Number(s.gpus_per_node), ring_degree: 1 }),
|
||||
validateTopology: (s, topology) => {
|
||||
const errors = [];
|
||||
const nodes = Number(s.nodes);
|
||||
const perNode = Number(s.gpus_per_node);
|
||||
const { tp_size: tp, ulysses_degree: ulysses, ring_degree: ring } = topology;
|
||||
if (nodes !== 1) errors.push("This picker covers single-node deployment only.");
|
||||
if (![1, 2, 4].includes(perNode)) errors.push("Select one, two, or four GPUs per node.");
|
||||
if (![tp, ulysses, ring].every((n) => [1, 2, 4].includes(n))) errors.push("TP, Ulysses and Ring must each be 1, 2, or 4.");
|
||||
if (nodes * perNode !== tp * ulysses * ring) errors.push(`World size ${nodes * perNode} must equal TP × Ulysses × Ring (${tp * ulysses * ring}).`);
|
||||
if (32 % (tp * ulysses) !== 0) errors.push("32 attention heads must be divisible by TP × Ulysses.");
|
||||
if (ring > 1 && effectiveAttention(s) === "sdpa") errors.push("Ring requires FlashAttention or SageAttention; Torch SDPA is unsupported.");
|
||||
if (s.precision?.startsWith("nvfp4_") && !["b200", "rtxpro6000", "rtx5090"].includes(s.hw)) errors.push("Native NVFP4 requires a Blackwell GPU. Select B200, RTX PRO 6000, or RTX 5090.");
|
||||
if (perNode === 1 && ["rtx5090", "rtx4090"].includes(s.hw) && s.placement === "resident") errors.push("The full resident pipeline exceeds this GPU's memory. Select CPU offload.");
|
||||
return errors;
|
||||
},
|
||||
},
|
||||
resolveDeployment: (s) => {
|
||||
const resource = config.commandBuilder.resource;
|
||||
const topology = s.topology_mode === "manual"
|
||||
? { tp_size: Number(s.tp_size), ulysses_degree: Number(s.ulysses_degree), ring_degree: Number(s.ring_degree) }
|
||||
: resource.autoTopology(s);
|
||||
const errors = resource.validateTopology(s, topology);
|
||||
const recipe = resource.verifiedRecipes.find((entry) => entry.hw === s.hw
|
||||
&& entry.nodes === Number(s.nodes) && entry.gpus_per_node === Number(s.gpus_per_node)
|
||||
&& entry.placement === s.placement && entry.tp_size === topology.tp_size
|
||||
&& entry.ulysses_degree === topology.ulysses_degree && entry.ring_degree === topology.ring_degree);
|
||||
const serveVerified = !!recipe && errors.length === 0 && s.encoder === "auto"
|
||||
&& recipe.attentions.includes(effectiveAttention(s)) && s.precision === "native"
|
||||
&& s.execution === "eager" && s.vae === "full";
|
||||
// Exact HTTP workloads from the validation matrix, not blanket quality coverage.
|
||||
const requestVerified = serveVerified
|
||||
&& ((["text", "edit"].includes(s.mode) && s.resolution === "1024" && Number(s.steps) === 40 && Number(s.outputs) === 1
|
||||
&& (["h200", "rtxpro6000"].includes(s.hw) || s.mode === "text" || s.background === "scene"))
|
||||
|| (s.hw === "h200" && s.background === "scene" && s.mode === "text" && s.resolution === "512" && Number(s.steps) === 4 && Number(s.outputs) === 2)
|
||||
|| (s.hw === "h200" && s.background === "scene" && s.mode === "multi" && s.resolution === "512" && Number(s.steps) === 4 && Number(s.outputs) === 1));
|
||||
const world = Number(s.nodes) * Number(s.gpus_per_node);
|
||||
const flags = ['--model-path "{{MODEL_PATH}}"', "--model-id Qwen-Image-2.1", `--num-gpus ${world}`];
|
||||
if (topology.tp_size > 1) flags.push(`--tp-size ${topology.tp_size}`);
|
||||
flags.push(`--ulysses-degree ${topology.ulysses_degree}`);
|
||||
if (topology.ring_degree > 1) flags.push(`--ring-degree ${topology.ring_degree}`);
|
||||
flags.push("--host {{HOST_IP}}", "--port {{PORT}}");
|
||||
const warnings = [];
|
||||
if (!serveVerified && !errors.length) warnings.push("This server combination has not completed an exact HTTP verification run.");
|
||||
if (!requestVerified && !errors.length) warnings.push("This request shape is outside the verified HTTP matrix.");
|
||||
return {
|
||||
match: { hw: s.hw }, nnodes: Number(s.nodes), verified: serveVerified, flags,
|
||||
builder: {
|
||||
topology,
|
||||
topologySummary: `TP ${topology.tp_size} · Ulysses ${topology.ulysses_degree} · Ring ${topology.ring_degree}`,
|
||||
errors, warnings,
|
||||
verification: {
|
||||
serve: errors.length ? "error" : serveVerified ? "verified" : "unverified",
|
||||
request: errors.length ? "error" : requestVerified ? "verified" : "unverified",
|
||||
},
|
||||
resolvedSettings: {
|
||||
attention: s.attention === "platform" ? `${platformAttention(s) === "sdpa" ? "Torch SDPA" : "FlashAttention"} (auto)`
|
||||
: sm120Hardware.includes(s.hw) && s.attention === "fa" ? "Torch SDPA (FA fallback)" : undefined,
|
||||
encoder: s.encoder === "auto" && world === 1 ? "Single GPU (auto)" : undefined,
|
||||
},
|
||||
},
|
||||
};
|
||||
},
|
||||
},
|
||||
|
||||
modelNames: { default: "Qwen-Image-2.1" },
|
||||
placeholders: {
|
||||
MODEL_PATH: { target: "command", label: "Authorized checkpoint directory", default: "/models/qwen-image-2.1" },
|
||||
FP8_DIT_PATH: { target: "command", label: "Serialized FP8 DiT directory", default: "/models/qwen-image-2.1-fp8/transformer" },
|
||||
FP8_ENCODER_PATH: { target: "command", label: "Serialized FP8 encoder directory", default: "/models/qwen-image-2.1-fp8/text_encoder" },
|
||||
GGUF_DIT_PATH: { target: "command", label: "GGUF DiT file", default: "/models/qwen-image-2.1-gguf/transformer-Q4_0.gguf" },
|
||||
GGUF_ENCODER_PATH: { target: "command", label: "GGUF encoder file", default: "/models/qwen-image-2.1-gguf/text_encoder-Q4_0.gguf" },
|
||||
NVFP4_DIT_PATH: { target: "command", label: "NVFP4 DiT directory", default: "/models/qwen-image-2.1-nvfp4/transformer" },
|
||||
NVFP4_ENCODER_PATH: { target: "command", label: "NVFP4 encoder directory", default: "/models/qwen-image-2.1-nvfp4/text_encoder" },
|
||||
HOST_IP: { target: "command", label: "Bind host", default: "0.0.0.0" },
|
||||
PORT: { target: "command", label: "Bind port", default: "30010" },
|
||||
CURL_HOST: { target: "curl", label: "Server host", default: "localhost" },
|
||||
CURL_PORT: { target: "curl", label: "Server port", default: "30010" },
|
||||
INPUT_IMAGE: { target: "curl", label: "First reference PNG (client path)", default: "/path/to/input.png" },
|
||||
SECOND_IMAGE: { target: "curl", label: "Second reference PNG (client path)", default: "/path/to/reference.png" },
|
||||
},
|
||||
curl: (s) => {
|
||||
const transparent = s.background === "transparent";
|
||||
const prompts = {
|
||||
text: transparent
|
||||
? "A single fluffy orange cat sitting, full body, isolated on a transparent background. A clean cutout with an alpha channel, transparent outside the cat, no floor, no shadow, no background."
|
||||
: "A capybara reading a book by candlelight",
|
||||
edit: transparent
|
||||
? "Change the orange fur of the cat to gray, keeping its pose, shape and fur detail unchanged. Preserve the transparent background and alpha channel. No floor, no shadow, no background."
|
||||
: "Change the red teapot to blue, keeping its shape, table, window, and lighting unchanged.",
|
||||
multi: transparent
|
||||
? "Combine the subjects from Picture 1 and Picture 2 into one composition on a transparent background. Preserve an alpha channel outside the subjects."
|
||||
: "Combine the subjects from Picture 1 and Picture 2 into one coherent scene, preserving their appearance.",
|
||||
};
|
||||
const request = {
|
||||
model: "{{MODEL_NAME}}", prompt: prompts[s.mode], n: Number(s.outputs),
|
||||
size: `${s.resolution}x${s.resolution}`, num_inference_steps: Number(s.steps),
|
||||
guidance_scale: 1, seed: 42, generator_device: "cpu",
|
||||
output_format: "png", response_format: "b64_json",
|
||||
background: transparent ? "transparent" : "auto",
|
||||
};
|
||||
if (s.mode === "text") {
|
||||
return `curl -sS --fail-with-body http://{{CURL_HOST}}:{{CURL_PORT}}/v1/images/generations \\
|
||||
-H 'Content-Type: application/json' \\
|
||||
-d '${JSON.stringify({ ...request, enable_cache_dit: false }, null, 2)}'`;
|
||||
}
|
||||
const fields = Object.entries(request).map(([key, value]) => ` --form-string '${key}=${value}'`);
|
||||
fields.push(' -F "image[]=@{{INPUT_IMAGE}};type=image/png"');
|
||||
if (s.mode === "multi") fields.push(' -F "image[]=@{{SECOND_IMAGE}};type=image/png"');
|
||||
return `curl -sS --fail-with-body http://{{CURL_HOST}}:{{CURL_PORT}}/v1/images/edits \\
|
||||
${fields.join(" \\\n")}`;
|
||||
},
|
||||
// The integration is installed from source; no published Docker image is verified.
|
||||
runModes: () => ["python"],
|
||||
showPlaygroundLink: false,
|
||||
cells: [],
|
||||
};
|
||||
|
||||
return config;
|
||||
})();
|
||||
@@ -25,6 +25,11 @@ export const DiffusionModelCatalog = ({ category }) => {
|
||||
],
|
||||
cookbook: "/cookbook/diffusion/Qwen-Image/Qwen-Image",
|
||||
},
|
||||
{
|
||||
name: "Qwen-Image 2.1",
|
||||
modelIds: ["Qwen/Qwen-Image-2.1"],
|
||||
cookbook: "/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1",
|
||||
},
|
||||
{
|
||||
name: "Qwen-Image Edit / Layered",
|
||||
modelIds: [
|
||||
|
||||
@@ -190,6 +190,7 @@ def _layernorm_modulate_kernel(
|
||||
FP8_MAX: tl.constexpr,
|
||||
STORE_BF16: tl.constexpr,
|
||||
QUANTIZE_FP8: tl.constexpr,
|
||||
HAS_SHIFT: tl.constexpr = True,
|
||||
):
|
||||
pid = tl.program_id(0).to(tl.int64)
|
||||
row_offs = pid * ROWS + tl.arange(0, ROWS)
|
||||
@@ -261,13 +262,15 @@ def _layernorm_modulate_kernel(
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
sh = tl.load(
|
||||
shift_ptr + batch[:, None] * scale_row_stride + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
one_plus = round_bf16_to_fp32(1.0 + sc)
|
||||
y = round_bf16_to_fp32(y * one_plus) + sh
|
||||
y = round_bf16_to_fp32(y * one_plus)
|
||||
if HAS_SHIFT:
|
||||
sh = tl.load(
|
||||
shift_ptr + batch[:, None] * scale_row_stride + cols[None, :],
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
y = y + sh
|
||||
if STORE_BF16:
|
||||
tl.store(y_ptr + row_base[:, None] + cols[None, :], y, mask=mask)
|
||||
if QUANTIZE_FP8:
|
||||
@@ -388,7 +391,7 @@ def _mod_row_stride(t: torch.Tensor, batch: int, hidden: int) -> int | None:
|
||||
|
||||
|
||||
def can_use_fused_layernorm_modulate(
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor | None
|
||||
) -> bool:
|
||||
if not (
|
||||
_is_bf16_cuda(x)
|
||||
@@ -398,29 +401,34 @@ def can_use_fused_layernorm_modulate(
|
||||
and x.shape[-1] % 4 == 0
|
||||
and x.shape[-1] <= 8192
|
||||
and _is_bf16_cuda(scale)
|
||||
and _is_bf16_cuda(shift)
|
||||
and scale.device == x.device
|
||||
and shift.device == x.device
|
||||
):
|
||||
return False
|
||||
batch, _, hidden = x.shape
|
||||
q = _mod_row_stride(scale, batch, hidden)
|
||||
if shift is None:
|
||||
return q is not None
|
||||
if not _is_bf16_cuda(shift) or shift.device != x.device:
|
||||
return False
|
||||
v = _mod_row_stride(shift, batch, hidden)
|
||||
return q is not None and v is not None and q == v
|
||||
|
||||
|
||||
def _fake_ln_modulate(
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor, eps: float
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor | None, eps: float
|
||||
) -> torch.Tensor:
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
def fused_layernorm_modulate_raw(
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor, eps: float
|
||||
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor | None, eps: float
|
||||
) -> torch.Tensor:
|
||||
"""``LN(x) * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)``, bit-exact
|
||||
vs the eager aten chain (LayerNorm without affine).
|
||||
|
||||
With ``shift=None``, omit the addition, preserving signed zeros in
|
||||
scale-only modulation.
|
||||
|
||||
Direct-call variant without the ``torch.ops`` dispatch (which costs tens
|
||||
of microseconds per call); use it on CPU-launch-bound eager hot paths
|
||||
(e.g. Sana), and the registered custom op under ``torch.compile``.
|
||||
@@ -449,6 +457,7 @@ def fused_layernorm_modulate_raw(
|
||||
FP8_MAX=fp8_max,
|
||||
STORE_BF16=True,
|
||||
QUANTIZE_FP8=False,
|
||||
HAS_SHIFT=shift is not None,
|
||||
# H200-tuned: 38.5us at (1, 4096, 4096) vs the 121.8us eager
|
||||
# chain, 14.3us at Sana's (2, 1024, 2240) vs 43.1us. ROWS=1 +
|
||||
# 4 warps triggers pathological Triton layout conversions in
|
||||
|
||||
@@ -145,6 +145,8 @@ tensor copy per residual site.
|
||||
| `try_fused_flux2_qkv_epilogue` | KDA (JIT CUDA) | bit-exact vs the selected BF16 chain | FLUX.2 QK RMSNorm + RoPE + joint QKV packing |
|
||||
| `try_fused_qwen_qkv_epilogue` | JIT CUDA | bit-exact vs the selected BF16 chain | Qwen-Image QK RMSNorm + RoPE + joint QKV writes; SM90+ |
|
||||
| `fused_rope_rotate_half_bitexact` | Triton | bit-exact (elementwise only) |
|
||||
| `fused_complex_rope` | Triton | preserves CUDA complex64 multiply rounding for contiguous BSHD inputs; Qwen-Image 2.1 verifies its first call against eager |
|
||||
| `rmsnorm_preserve_reduction` | Triton + aten | preserves the FP32 mean reduction and cast-before-weight rounding; fuses only pointwise work for contiguous FP16/BF16 inputs; Qwen-Image 2.1 verifies its first call |
|
||||
| `fused_interleaved_rope_fp64` | JIT CUDA | bit-exact vs paired SANA-Video fp64 RoPE |
|
||||
| `fused_inplace_helios_qk_rope` | JIT CUDA | bit-exact paired in-place RoPE for Helios' transposed frequency layout |
|
||||
| `ltx2_qknorm_split_rope_cuda` | KDA (JIT CUDA) | close; **validated on B200** |
|
||||
|
||||
@@ -280,6 +280,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = (
|
||||
_CUDA,
|
||||
"Paired in-place Helios transposed Q/K RoPE.",
|
||||
),
|
||||
(
|
||||
"diffusion.complex_rope",
|
||||
KernelBackend.TRITON,
|
||||
"rope.complex_rope_triton:fused_complex_rope",
|
||||
_CUDA,
|
||||
"Paired RoPE preserving PyTorch complex64 multiplication rounding.",
|
||||
),
|
||||
(
|
||||
"diffusion.hunyuan_qkv_rope_pack",
|
||||
KernelBackend.TRITON,
|
||||
@@ -287,6 +294,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = (
|
||||
_CUDA,
|
||||
"HunyuanVideo QKV pack + RoPE.",
|
||||
),
|
||||
(
|
||||
"diffusion.rmsnorm_preserve_reduction",
|
||||
KernelBackend.TRITON,
|
||||
"norm.rmsnorm_preserve_reduction:rmsnorm_preserve_reduction",
|
||||
_CUDA,
|
||||
"Cast-before-weight RMSNorm preserving the native FP32 mean reduction.",
|
||||
),
|
||||
(
|
||||
"diffusion.silu_mul",
|
||||
KernelBackend.TRITON,
|
||||
@@ -536,6 +550,8 @@ _EXPORTS: dict[str, str] = {
|
||||
"try_fused_bias_mul_add": "sglang.kernels.kda_kernels.norm_scale_shift_jit",
|
||||
"try_fused_bias_scale_residual_norm_scale_shift": "sglang.kernels.kda_kernels.norm_scale_shift_jit",
|
||||
"triton_one_pass_rms_norm": "norm.rmsnorm_onepass_triton",
|
||||
"can_use_rmsnorm_preserve_reduction": "norm.rmsnorm_preserve_reduction",
|
||||
"rmsnorm_preserve_reduction": "norm.rmsnorm_preserve_reduction",
|
||||
"can_use_fused_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact",
|
||||
"can_use_fused_scale_residual_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact",
|
||||
"fused_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact",
|
||||
@@ -589,6 +605,8 @@ _EXPORTS: dict[str, str] = {
|
||||
"can_use_helios_qk_rope": "rope.helios_qk_rope_jit",
|
||||
"fused_inplace_helios_qk_rope": "rope.helios_qk_rope_jit",
|
||||
"apply_rotary_embedding": "rope.rotary_triton",
|
||||
"can_use_fused_complex_rope": "rope.complex_rope_triton",
|
||||
"fused_complex_rope": "rope.complex_rope_triton",
|
||||
# Tensor layout transformations fused with downstream quantization
|
||||
"try_flux2_token_cat_fp8": "sglang.kernels.kda_kernels.flux2_token_cat_fp8_triton",
|
||||
# Activation-function fusions
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fuse channel-first RMSNorm pointwise work, preserving native FP32 L2 norm."""
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _channel_rmsnorm_finish_kernel(
|
||||
x_ptr,
|
||||
norm_ptr,
|
||||
weight_ptr,
|
||||
out_ptr,
|
||||
N: tl.constexpr,
|
||||
CHANNELS: tl.constexpr,
|
||||
SPATIAL: tl.constexpr,
|
||||
SCALE: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
|
||||
mask = index < N
|
||||
channel = index // SPATIAL % CHANNELS
|
||||
norm_index = index // (CHANNELS * SPATIAL) * SPATIAL + index % SPATIAL
|
||||
value = tl.load(x_ptr + index, mask, 0).to(tl.float32)
|
||||
norm = tl.maximum(tl.load(norm_ptr + norm_index, mask, 1), 1.0e-12)
|
||||
weight = tl.load(weight_ptr + channel, mask, 0).to(tl.float32)
|
||||
# F.normalize divides in FP32, then the original module rounds both the
|
||||
# normalized activation and its scale multiply before applying gamma.
|
||||
value = tl.div_rn(value, norm).to(x_ptr.dtype.element_ty).to(tl.float32)
|
||||
value = (value * SCALE).to(x_ptr.dtype.element_ty).to(tl.float32)
|
||||
value = (value * weight).to(x_ptr.dtype.element_ty).to(tl.float32)
|
||||
tl.store(out_ptr + index, value + 0.0, mask)
|
||||
|
||||
|
||||
def can_use_channel_rmsnorm(x, weight):
|
||||
return (
|
||||
x.is_cuda
|
||||
and torch.version.hip is None
|
||||
and x.dtype in (torch.bfloat16, torch.float16)
|
||||
and x.ndim in (4, 5)
|
||||
and x.numel() > 0
|
||||
and x.is_contiguous()
|
||||
and weight.device == x.device
|
||||
and weight.dtype == x.dtype
|
||||
and weight.shape == (x.shape[1],) + (1,) * (x.ndim - 2)
|
||||
and weight.is_contiguous()
|
||||
)
|
||||
|
||||
|
||||
def _fake_channel_rmsnorm(x, weight, scale):
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="channel_rmsnorm_preserve_reduction",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_channel_rmsnorm,
|
||||
)
|
||||
def channel_rmsnorm_preserve_reduction(
|
||||
x: torch.Tensor, weight: torch.Tensor, scale: float
|
||||
) -> torch.Tensor:
|
||||
assert can_use_channel_rmsnorm(x, weight)
|
||||
# Keep F.normalize's input dtype, shape and native reduction dispatch.
|
||||
norm = x.float().norm(p=2, dim=1, keepdim=True)
|
||||
out = torch.empty_like(x)
|
||||
with torch.cuda.device(x.device):
|
||||
_channel_rmsnorm_finish_kernel[(triton.cdiv(x.numel(), 512),)](
|
||||
x,
|
||||
norm,
|
||||
weight,
|
||||
out,
|
||||
x.numel(),
|
||||
x.shape[1],
|
||||
x.numel() // (x.shape[0] * x.shape[1]),
|
||||
scale,
|
||||
512,
|
||||
enable_fp_fusion=False,
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,92 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fuse RMSNorm pointwise work while retaining the native FP32 mean reduction.
|
||||
|
||||
Matches ``weight * (x.float() * rsqrt(mean(x.float()**2) + eps)).to(x.dtype)``
|
||||
for contiguous FP16/BF16 inputs. The FP32 square buffer has the original shape,
|
||||
so aten selects the same reduction as the eager chain. Verified at head width
|
||||
128, including 131072 rows; callers verify their first dispatch before reuse.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _square_fp32_kernel(x_ptr, square_ptr, N: tl.constexpr, BLOCK: tl.constexpr):
|
||||
index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
|
||||
value = tl.load(x_ptr + index, index < N, 0).to(tl.float32)
|
||||
tl.store(square_ptr + index, value * value, index < N)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _rmsnorm_finish_kernel(
|
||||
x_ptr,
|
||||
variance_ptr,
|
||||
weight_ptr,
|
||||
out_ptr,
|
||||
N: tl.constexpr,
|
||||
DIM: tl.constexpr,
|
||||
EPS: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
index = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
|
||||
mask = index < N
|
||||
value = tl.load(x_ptr + index, mask, 0).to(tl.float32)
|
||||
variance = tl.load(variance_ptr + index // DIM, mask, 0)
|
||||
weight = tl.load(weight_ptr + index % DIM, mask, 0).to(tl.float32)
|
||||
# eager rounds the normalized activation before multiplying the weight
|
||||
normalized = (value * tl.rsqrt(variance + EPS)).to(x_ptr.dtype.element_ty)
|
||||
tl.store(out_ptr + index, normalized.to(tl.float32) * weight, mask)
|
||||
|
||||
|
||||
def can_use_rmsnorm_preserve_reduction(x: torch.Tensor, weight: torch.Tensor) -> bool:
|
||||
return (
|
||||
x.is_cuda
|
||||
and torch.version.hip is None
|
||||
and x.dtype in (torch.float16, torch.bfloat16)
|
||||
and x.ndim >= 2
|
||||
and x.numel() > 0
|
||||
and x.is_contiguous()
|
||||
and weight.device == x.device
|
||||
and weight.dtype == x.dtype
|
||||
and weight.shape == (x.shape[-1],)
|
||||
and weight.is_contiguous()
|
||||
)
|
||||
|
||||
|
||||
def _fake_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="rmsnorm_preserve_reduction",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_rmsnorm,
|
||||
)
|
||||
def rmsnorm_preserve_reduction(
|
||||
x: torch.Tensor, weight: torch.Tensor, eps: float
|
||||
) -> torch.Tensor:
|
||||
"""Preserve aten's mean and cast-before-weight semantics without residuals."""
|
||||
assert can_use_rmsnorm_preserve_reduction(x, weight)
|
||||
squares = torch.empty_like(x, dtype=torch.float32)
|
||||
out = torch.empty_like(x)
|
||||
with torch.cuda.device(x.device):
|
||||
_square_fp32_kernel[(triton.cdiv(x.numel(), 1024),)](
|
||||
x, squares, x.numel(), 1024
|
||||
)
|
||||
variance = squares.mean(dim=-1, keepdim=True)
|
||||
_rmsnorm_finish_kernel[(triton.cdiv(x.numel(), 512),)](
|
||||
x,
|
||||
variance,
|
||||
weight,
|
||||
out,
|
||||
x.numel(),
|
||||
x.shape[-1],
|
||||
eps,
|
||||
512,
|
||||
enable_fp_fusion=False,
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,96 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Paired RoPE preserving PyTorch CUDA complex64 multiplication rounding."""
|
||||
|
||||
from functools import lru_cache
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _complex_rope_kernel(
|
||||
x_ptr,
|
||||
rope_ptr,
|
||||
out_ptr,
|
||||
pairs,
|
||||
SEQ: tl.constexpr,
|
||||
HEADS: tl.constexpr,
|
||||
DIM: tl.constexpr,
|
||||
FUSE_REAL_SIN: tl.constexpr,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
pair = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
|
||||
mask = pair < pairs
|
||||
token = pair // (HEADS * (DIM // 2)) % SEQ
|
||||
column = pair % (DIM // 2)
|
||||
real = tl.load(x_ptr + 2 * pair, mask, 0).to(tl.float32)
|
||||
imag = tl.load(x_ptr + 2 * pair + 1, mask, 0).to(tl.float32)
|
||||
cos = tl.load(rope_ptr + token * DIM + 2 * column, mask, 0)
|
||||
sin = tl.load(rope_ptr + token * DIM + 2 * column + 1, mask, 0)
|
||||
# CUDA builds differ in which imaginary product is contracted into the FMA
|
||||
out_real = tl.fma(real, cos, -imag * sin)
|
||||
if FUSE_REAL_SIN:
|
||||
out_imag = tl.fma(real, sin, imag * cos)
|
||||
else:
|
||||
out_imag = tl.fma(imag, cos, real * sin)
|
||||
tl.store(out_ptr + 2 * pair, out_real, mask)
|
||||
tl.store(out_ptr + 2 * pair + 1, out_imag, mask)
|
||||
|
||||
|
||||
@lru_cache
|
||||
def _fuse_real_sin(device: torch.device) -> bool:
|
||||
# cancellation distinguishes the two orders without depending on the GPU name
|
||||
values = torch.tensor(
|
||||
[[1 + 2**-23, -1], [1, 1 - 2**-24]], device=device, dtype=torch.float32
|
||||
)
|
||||
z = torch.view_as_complex(values)
|
||||
return (z[0] * z[1]).imag.item() != 0
|
||||
|
||||
|
||||
def can_use_fused_complex_rope(x: torch.Tensor, rope: torch.Tensor) -> bool:
|
||||
return (
|
||||
x.is_cuda
|
||||
and torch.version.hip is None
|
||||
and x.dtype in (torch.float16, torch.bfloat16, torch.float32)
|
||||
and x.ndim == 4
|
||||
and x.numel() > 0
|
||||
and x.shape[-1] % 2 == 0
|
||||
and x.is_contiguous()
|
||||
and rope.dtype == torch.complex64
|
||||
and rope.device == x.device
|
||||
and rope.shape == (x.shape[1], x.shape[-1] // 2)
|
||||
and rope.is_contiguous()
|
||||
)
|
||||
|
||||
|
||||
def _fake_complex_rope(x: torch.Tensor, rope: torch.Tensor) -> torch.Tensor:
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="fused_complex_rope",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_complex_rope,
|
||||
)
|
||||
def fused_complex_rope(x: torch.Tensor, rope: torch.Tensor) -> torch.Tensor:
|
||||
"""Rotate contiguous BSHD activations with a shared S×(D/2) complex cache."""
|
||||
assert can_use_fused_complex_rope(x, rope)
|
||||
out = torch.empty_like(x)
|
||||
pairs = x.numel() // 2
|
||||
with torch.cuda.device(x.device):
|
||||
_complex_rope_kernel[(triton.cdiv(pairs, 256),)](
|
||||
x,
|
||||
torch.view_as_real(rope),
|
||||
out,
|
||||
pairs,
|
||||
x.shape[1],
|
||||
x.shape[2],
|
||||
x.shape[3],
|
||||
_fuse_real_sin(x.device),
|
||||
256,
|
||||
enable_fp_fusion=False,
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,123 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Write normalized/rotated K and unmodified V into their final prefix buffers."""
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.kernels.ops.diffusion.rope.complex_rope_triton import _fuse_real_sin
|
||||
from sglang.kernels.ops.diffusion.rope.qknorm_complex_rope_triton import (
|
||||
_qknorm_complex_rope_rows,
|
||||
can_use_qknorm_complex_rope,
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _qknorm_complex_rope_kv_kernel(
|
||||
k_ptr,
|
||||
weight_ptr,
|
||||
rope_ptr,
|
||||
v_ptr,
|
||||
kp_ptr,
|
||||
vp_ptr,
|
||||
kout_ptr,
|
||||
vout_ptr,
|
||||
ROWS: tl.constexpr,
|
||||
SEQ: tl.constexpr,
|
||||
HEADS: tl.constexpr,
|
||||
PREFIX: tl.constexpr,
|
||||
BATCH: tl.constexpr,
|
||||
EPS: tl.constexpr,
|
||||
FUSE_REAL_SIN: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
if pid < tl.cdiv(ROWS, 4):
|
||||
row = pid * 4 + tl.arange(0, 4)
|
||||
column = tl.arange(0, 128)
|
||||
key = _qknorm_complex_rope_rows(
|
||||
k_ptr, weight_ptr, rope_ptr, row, ROWS, SEQ, HEADS, EPS, FUSE_REAL_SIN
|
||||
)
|
||||
out_row = row + (row // (SEQ * HEADS) + 1) * PREFIX * HEADS
|
||||
output_index = out_row[:, None] * 128 + column[None, :]
|
||||
mask = row[:, None] < ROWS
|
||||
tl.store(kout_ptr + output_index, key, mask)
|
||||
value = tl.load(v_ptr + row[:, None] * 128 + column[None, :], mask, 0)
|
||||
tl.store(vout_ptr + output_index, value, mask)
|
||||
else:
|
||||
index = (pid - tl.cdiv(ROWS, 4)) * 1024 + tl.arange(0, 1024)
|
||||
prefix_mask = index < BATCH * PREFIX * HEADS * 128
|
||||
prefix_index = index + (index // (PREFIX * HEADS * 128)) * SEQ * HEADS * 128
|
||||
prefix_key = tl.load(kp_ptr + index, prefix_mask, 0)
|
||||
prefix_value = tl.load(vp_ptr + index, prefix_mask, 0)
|
||||
tl.store(kout_ptr + prefix_index, prefix_key, prefix_mask)
|
||||
tl.store(vout_ptr + prefix_index, prefix_value, prefix_mask)
|
||||
|
||||
|
||||
def can_use_qknorm_complex_rope_kv(k, weight, rope, v, k_prefix, v_prefix):
|
||||
return (
|
||||
can_use_qknorm_complex_rope(k, weight, rope)
|
||||
and v.shape == k.shape
|
||||
and k_prefix.ndim == 4
|
||||
and k_prefix.shape[0] == k.shape[0]
|
||||
and k_prefix.shape[1] > 0
|
||||
and k_prefix.shape[2:] == k.shape[2:]
|
||||
and v_prefix.shape == k_prefix.shape
|
||||
and all(
|
||||
x.device == k.device and x.dtype == k.dtype and x.is_contiguous()
|
||||
for x in (v, k_prefix, v_prefix)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _fake_qknorm_complex_rope_kv(k, weight, rope, v, k_prefix, v_prefix, eps):
|
||||
shape = (k.shape[0], k_prefix.shape[1] + k.shape[1], *k.shape[2:])
|
||||
return k.new_empty(shape), v.new_empty(shape)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="qknorm_complex_rope_kv",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_qknorm_complex_rope_kv,
|
||||
)
|
||||
def qknorm_complex_rope_kv(
|
||||
k: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
rope: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
k_prefix: torch.Tensor,
|
||||
v_prefix: torch.Tensor,
|
||||
eps: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
assert can_use_qknorm_complex_rope_kv(k, weight, rope, v, k_prefix, v_prefix)
|
||||
kout, vout = _fake_qknorm_complex_rope_kv(
|
||||
k, weight, rope, v, k_prefix, v_prefix, eps
|
||||
)
|
||||
batch, seq, heads, dim = k.shape
|
||||
prefix = k_prefix.shape[1]
|
||||
with torch.cuda.device(k.device):
|
||||
_qknorm_complex_rope_kv_kernel[
|
||||
(
|
||||
triton.cdiv(batch * seq * heads, 4)
|
||||
+ triton.cdiv(batch * prefix * heads * dim, 1024),
|
||||
)
|
||||
](
|
||||
k,
|
||||
weight,
|
||||
torch.view_as_real(rope),
|
||||
v,
|
||||
k_prefix,
|
||||
v_prefix,
|
||||
kout,
|
||||
vout,
|
||||
batch * seq * heads,
|
||||
seq,
|
||||
heads,
|
||||
prefix,
|
||||
batch,
|
||||
eps,
|
||||
_fuse_real_sin(k.device),
|
||||
num_warps=4,
|
||||
enable_fp_fusion=False,
|
||||
)
|
||||
return kout, vout
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fuse 128-wide RMSNorm and complex RoPE with native rounding boundaries."""
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.kernels.ops.diffusion.norm.rmsnorm_preserve_reduction import (
|
||||
can_use_rmsnorm_preserve_reduction,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.rope.complex_rope_triton import (
|
||||
_fuse_real_sin,
|
||||
can_use_fused_complex_rope,
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _qknorm_complex_rope_rows(
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
rope_ptr,
|
||||
row,
|
||||
ROWS: tl.constexpr,
|
||||
SEQ: tl.constexpr,
|
||||
HEADS: tl.constexpr,
|
||||
EPS: tl.constexpr,
|
||||
FUSE_REAL_SIN: tl.constexpr,
|
||||
):
|
||||
# Four rows / four warps gives each lane four consecutive components.
|
||||
# Match aten's vectorized 128-wide FP32 mean: combine four components
|
||||
# left-to-right, then reduce 32 lanes with decreasing shuffle offsets.
|
||||
# Increasing rows per warp changes this order and is not bit-exact.
|
||||
column = tl.arange(0, 128)
|
||||
mask = row[:, None] < ROWS
|
||||
value = tl.load(x_ptr + row[:, None] * 128 + column[None, :], mask, 0).to(
|
||||
tl.float32
|
||||
)
|
||||
square = tl.reshape(value * value, (4, 32, 2, 2))
|
||||
even, odd = tl.split(square)
|
||||
a, c = tl.split(even)
|
||||
b, d = tl.split(odd)
|
||||
variance = tl.sum(((a + b) + c) + d, 1) * (1.0 / 128)
|
||||
inv = tl.rsqrt(variance + EPS)
|
||||
weight = tl.load(weight_ptr + column).to(tl.float32)
|
||||
value = (value * inv[:, None]).to(x_ptr.dtype.element_ty).to(tl.float32)
|
||||
value = (value * weight[None, :]).to(x_ptr.dtype.element_ty).to(tl.float32)
|
||||
real, imag = tl.split(tl.reshape(value, (4, 64, 2)))
|
||||
token = row // HEADS % SEQ
|
||||
rotation = tl.load(rope_ptr + token[:, None] * 128 + column[None, :], mask, 0)
|
||||
cos, sin = tl.split(tl.reshape(rotation, (4, 64, 2)))
|
||||
out_real = tl.fma(real, cos, -imag * sin)
|
||||
if FUSE_REAL_SIN:
|
||||
out_imag = tl.fma(real, sin, imag * cos)
|
||||
else:
|
||||
out_imag = tl.fma(imag, cos, real * sin)
|
||||
return tl.reshape(tl.join(out_real, out_imag), (4, 128))
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _qknorm_complex_rope_onepass_kernel(
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
rope_ptr,
|
||||
out_ptr,
|
||||
ROWS: tl.constexpr,
|
||||
SEQ: tl.constexpr,
|
||||
HEADS: tl.constexpr,
|
||||
EPS: tl.constexpr,
|
||||
FUSE_REAL_SIN: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0) * 4 + tl.arange(0, 4)
|
||||
out = _qknorm_complex_rope_rows(
|
||||
x_ptr, weight_ptr, rope_ptr, row, ROWS, SEQ, HEADS, EPS, FUSE_REAL_SIN
|
||||
)
|
||||
tl.store(
|
||||
out_ptr + row[:, None] * 128 + tl.arange(0, 128)[None, :],
|
||||
out,
|
||||
row[:, None] < ROWS,
|
||||
)
|
||||
|
||||
|
||||
def can_use_qknorm_complex_rope(x, weight, rope):
|
||||
return (
|
||||
can_use_rmsnorm_preserve_reduction(x, weight)
|
||||
and can_use_fused_complex_rope(x, rope)
|
||||
and x.shape[-1] == 128
|
||||
)
|
||||
|
||||
|
||||
def _fake_qknorm_complex_rope(x, weight, rope, eps):
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="qknorm_complex_rope",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_qknorm_complex_rope,
|
||||
)
|
||||
def qknorm_complex_rope(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
rope: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
assert can_use_qknorm_complex_rope(x, weight, rope)
|
||||
out = torch.empty_like(x)
|
||||
with torch.cuda.device(x.device):
|
||||
_qknorm_complex_rope_onepass_kernel[(triton.cdiv(x.numel() // 128, 4),)](
|
||||
x,
|
||||
weight,
|
||||
torch.view_as_real(rope),
|
||||
out,
|
||||
x.numel() // 128,
|
||||
x.shape[1],
|
||||
x.shape[2],
|
||||
eps,
|
||||
_fuse_real_sin(x.device),
|
||||
num_warps=4,
|
||||
enable_fp_fusion=False,
|
||||
)
|
||||
return out
|
||||
@@ -9,7 +9,7 @@ SGLang diffusion features an end-to-end unified pipeline for accelerating diffus
|
||||
## Key Features
|
||||
|
||||
SGLang Diffusion has the following features:
|
||||
- Broad model support: Wan, FastWan, FLUX, Qwen-Image, LongCat-Image, Z-Image, Ideogram 4, Krea-2, Cosmos3, LTX-2/LTX-2.3/LTX-2.5, MiniMax-H3, FastH3, VDN-H3, LingBot Video MoE, LingBot World, SANA-Video/SANA-WM, JoyEcho, MOVA, GLM-Image, ERNIE-Image, Hunyuan3D, and more
|
||||
- Broad model support: Wan, FastWan, FLUX, Qwen-Image / Qwen-Image 2.1, LongCat-Image, Z-Image, Ideogram 4, Krea-2, Cosmos3, LTX-2/LTX-2.3/LTX-2.5, MiniMax-H3, FastH3, VDN-H3, LingBot Video MoE, LingBot World, SANA-Video/SANA-WM, JoyEcho, MOVA, GLM-Image, ERNIE-Image, Hunyuan3D, and more
|
||||
- Fast inference speed: empowered by optimized `sgl-kernel` kernels, scheduler/runtime improvements, caching acceleration, and native diffusion hot-path optimizations
|
||||
- Ease of use: OpenAI-compatible api, CLI, and python sdk support
|
||||
- Multi-platform support:
|
||||
@@ -77,6 +77,24 @@ sglang generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--save-output
|
||||
```
|
||||
|
||||
### Qwen-Image 2.1
|
||||
|
||||
The native `QwenImage21Pipeline` supports text-to-image and reference-image
|
||||
conditioning with Qwen3-VL, a single-stream block-causal DiT, and the 64-channel
|
||||
VAE. Use an authorized checkpoint directory:
|
||||
|
||||
```bash
|
||||
sglang generate --model-path /models/qwen-image-2.1 --model-id Qwen-Image-2.1 \
|
||||
--prompt "A capybara reading a book by candlelight" \
|
||||
--height 1024 --width 1024 --num-inference-steps 40 --guidance-scale 1 \
|
||||
--seed 0 --save-output
|
||||
```
|
||||
|
||||
Add `--image-path /path/to/input.png` for editing. Dimensions must be multiples
|
||||
of 32. Full-checkpoint generation and editing have been tested on H200; see the
|
||||
[model cookbook](../../../docs/cookbook/diffusion/Qwen-Image/Qwen-Image-2.1.mdx)
|
||||
for component requirements and optimization boundaries.
|
||||
|
||||
### Component residency
|
||||
|
||||
Use `--component-residency COMPONENT=MODE` to choose one runtime mode for each
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21ArchConfig(DiTArchConfig):
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = 64
|
||||
num_layers: int = 32
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 32
|
||||
context_in_dim: int = 4096
|
||||
mlp_ratio: int = 3
|
||||
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
|
||||
eps: float = 1e-6
|
||||
causal_condition: bool = True
|
||||
causal_block: bool = True
|
||||
lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {r"^transformer\.": ""}
|
||||
)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
self.num_channels_latents = self.in_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21DitConfig(DiTConfig):
|
||||
arch_config: QwenImage21ArchConfig = field(default_factory=QwenImage21ArchConfig)
|
||||
prefix: str = "qwenimage21"
|
||||
@@ -0,0 +1,167 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from sglang.multimodal_gen.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21VAEArchConfig(VAEArchConfig):
|
||||
base_dim: int = 96
|
||||
decoder_base_dim: int = 144
|
||||
z_dim: int = 64
|
||||
dim_mult: tuple = (1, 2, 4, 8, 8)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple = ()
|
||||
temperal_downsample: tuple = (False, True, True, True)
|
||||
dropout: float = 0.0
|
||||
latents_mean: tuple = (
|
||||
0.5126,
|
||||
0.7721,
|
||||
-0.0631,
|
||||
1.3506,
|
||||
-0.7855,
|
||||
-2.1025,
|
||||
-0.3458,
|
||||
1.3722,
|
||||
1.8873,
|
||||
-1.7177,
|
||||
-0.651,
|
||||
0.2732,
|
||||
0.7562,
|
||||
-0.6163,
|
||||
-1.0277,
|
||||
3.8363,
|
||||
2.021,
|
||||
0.0472,
|
||||
0.932,
|
||||
2.0087,
|
||||
2.4954,
|
||||
-0.1391,
|
||||
-1.4249,
|
||||
1.8464,
|
||||
-0.5236,
|
||||
1.2826,
|
||||
3.7046,
|
||||
-1.3035,
|
||||
2.7286,
|
||||
-1.4518,
|
||||
-1.9036,
|
||||
-1.9955,
|
||||
-0.0342,
|
||||
-1.0265,
|
||||
-0.7636,
|
||||
3.0555,
|
||||
0.0746,
|
||||
-3.0751,
|
||||
-0.1076,
|
||||
1.7376,
|
||||
-1.0914,
|
||||
-1.9435,
|
||||
-0.2784,
|
||||
-1.368,
|
||||
0.4809,
|
||||
-0.4433,
|
||||
0.3764,
|
||||
0.5729,
|
||||
-2.0595,
|
||||
1.096,
|
||||
-1.326,
|
||||
-2.0211,
|
||||
-5.0179,
|
||||
0.5275,
|
||||
4.0162,
|
||||
1.8505,
|
||||
0.3026,
|
||||
1.9373,
|
||||
1.4937,
|
||||
0.2632,
|
||||
0.5547,
|
||||
-1.7121,
|
||||
-0.1562,
|
||||
0.0304,
|
||||
)
|
||||
latents_std: tuple = (
|
||||
3.2001,
|
||||
3.2936,
|
||||
3.4321,
|
||||
3.0091,
|
||||
3.1061,
|
||||
4.0379,
|
||||
4.0705,
|
||||
3.791,
|
||||
3.0785,
|
||||
3.65,
|
||||
3.9308,
|
||||
3.0904,
|
||||
2.8778,
|
||||
3.7675,
|
||||
3.732,
|
||||
5.0756,
|
||||
3.2864,
|
||||
4.0397,
|
||||
3.1317,
|
||||
4.0443,
|
||||
2.9249,
|
||||
3.9454,
|
||||
3.0988,
|
||||
4.2489,
|
||||
3.4896,
|
||||
3.8513,
|
||||
3.9323,
|
||||
3.4719,
|
||||
3.7498,
|
||||
4.283,
|
||||
3.5694,
|
||||
4.2467,
|
||||
3.9037,
|
||||
3.2947,
|
||||
5.077,
|
||||
3.5075,
|
||||
3.27,
|
||||
3.4767,
|
||||
2.8063,
|
||||
5.1125,
|
||||
3.5327,
|
||||
4.7833,
|
||||
3.1286,
|
||||
4.1819,
|
||||
3.8527,
|
||||
3.8312,
|
||||
3.5605,
|
||||
4.3875,
|
||||
3.9624,
|
||||
4.0168,
|
||||
3.5643,
|
||||
4.055,
|
||||
5.5614,
|
||||
4.2963,
|
||||
4.408,
|
||||
3.4959,
|
||||
3.8747,
|
||||
3.7608,
|
||||
3.5735,
|
||||
3.149,
|
||||
3.7662,
|
||||
3.6746,
|
||||
3.4563,
|
||||
3.8161,
|
||||
)
|
||||
is_residual: bool = True
|
||||
in_channels: int = 4
|
||||
out_channels: int = 4
|
||||
patch_size: int | None = None
|
||||
scale_factor_temporal: int = 8
|
||||
scale_factor_spatial: int = 16
|
||||
spatial_compression_ratio: int = 16
|
||||
temporal_compression_ratio: int = 1
|
||||
vae_scale_factor: int = 16
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21VAEConfig(VAEConfig):
|
||||
arch_config: QwenImage21VAEArchConfig = field(
|
||||
default_factory=QwenImage21VAEArchConfig
|
||||
)
|
||||
use_tiling: bool = False
|
||||
parallel_decode_mode: str = "tiled"
|
||||
use_temporal_tiling: bool = False
|
||||
@@ -0,0 +1,84 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.qwenimage21 import QwenImage21DitConfig
|
||||
from sglang.multimodal_gen.configs.models.encoders.qwen3vl import Qwen3VLConfig
|
||||
from sglang.multimodal_gen.configs.models.vaes.qwenimage21 import QwenImage21VAEConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.base import (
|
||||
ImagePipelineConfig,
|
||||
ModelTaskType,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21PipelineConfig(ImagePipelineConfig):
|
||||
native_only_components: tuple[str, ...] = ("transformer", "text_encoder", "vae")
|
||||
task_type: ModelTaskType = ModelTaskType.TI2I
|
||||
should_use_guidance: bool = False
|
||||
enable_autocast: bool = False
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
vae_precision: str = "bf16"
|
||||
generator_device: str = "cpu"
|
||||
dit_config: QwenImage21DitConfig = field(default_factory=QwenImage21DitConfig)
|
||||
vae_config: QwenImage21VAEConfig = field(default_factory=QwenImage21VAEConfig)
|
||||
text_encoder_configs: tuple = field(default_factory=lambda: (Qwen3VLConfig(),))
|
||||
text_encoder_precisions: tuple[str, ...] = ("bf16",)
|
||||
|
||||
def prepare_sigmas(self, sigmas, num_inference_steps):
|
||||
return self._prepare_sigmas(sigmas, num_inference_steps)
|
||||
|
||||
def get_classifier_free_guidance_scale(self, batch, guidance_scale):
|
||||
return (
|
||||
batch.true_cfg_scale if batch.true_cfg_scale is not None else guidance_scale
|
||||
)
|
||||
|
||||
def prepare_latent_shape(self, batch, batch_size, num_frames):
|
||||
return (
|
||||
batch_size,
|
||||
1,
|
||||
self.dit_config.in_channels,
|
||||
batch.height // 16,
|
||||
batch.width // 16,
|
||||
)
|
||||
|
||||
def maybe_pack_latents(self, latents, batch_size, batch):
|
||||
return latents.reshape(batch_size, self.dit_config.in_channels, -1).transpose(
|
||||
1, 2
|
||||
)
|
||||
|
||||
def shard_latents_for_sp(self, batch, latents):
|
||||
# the DiT shards only the target stream; its condition prefix stays replicated
|
||||
return latents, False
|
||||
|
||||
def gather_latents_for_sp(self, latents, batch=None):
|
||||
return latents
|
||||
|
||||
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||
return batch.extra["qwen21_positive"]
|
||||
|
||||
def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype=None):
|
||||
return batch.extra["qwen21_negative"]
|
||||
|
||||
def post_denoising_loop(self, latents, batch):
|
||||
# decode consumes only target latents, not the condition prefix or its KV cache
|
||||
batch.extra.pop("qwen21_positive", None)
|
||||
batch.extra.pop("qwen21_negative", None)
|
||||
return latents.transpose(1, 2).reshape(
|
||||
latents.shape[0], -1, 1, batch.height // 16, batch.width // 16
|
||||
)
|
||||
|
||||
def get_decode_scale_and_shift(self, device, dtype, vae):
|
||||
ac = self.vae_config.arch_config
|
||||
mean = torch.tensor(ac.latents_mean, device=device, dtype=dtype).view(
|
||||
1, ac.z_dim, 1, 1, 1
|
||||
)
|
||||
std = torch.tensor(ac.latents_std, device=device, dtype=dtype).view(
|
||||
1, ac.z_dim, 1, 1, 1
|
||||
)
|
||||
return std.reciprocal(), mean
|
||||
|
||||
def preprocess_condition_image(self, image, **kwargs):
|
||||
return image
|
||||
@@ -0,0 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from typing import ClassVar
|
||||
|
||||
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||
|
||||
|
||||
@dataclass
|
||||
class QwenImage21SamplingParams(SamplingParams):
|
||||
_default_height: ClassVar[int] = 1024
|
||||
_default_width: ClassVar[int] = 1024
|
||||
num_frames: int = 1
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 40
|
||||
negative_prompt: str | None = None
|
||||
@@ -99,6 +99,9 @@ from sglang.multimodal_gen.configs.pipeline_configs.qwen_image import (
|
||||
QwenImageLayeredPipelineConfig,
|
||||
QwenImagePipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import (
|
||||
QwenImage21PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.sana import SanaPipelineConfig
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.sana_video import (
|
||||
SanaVideoPipelineConfig,
|
||||
@@ -185,6 +188,7 @@ from sglang.multimodal_gen.configs.sample.qwenimage import (
|
||||
QwenImageLayeredSamplingParams,
|
||||
QwenImageSamplingParams,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.sample.qwenimage21 import QwenImage21SamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.sana import SanaSamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.sana_video import SanaVideoSamplingParams
|
||||
from sglang.multimodal_gen.configs.sample.sana_wm import SanaWMSamplingParams
|
||||
@@ -1131,6 +1135,12 @@ def _register_configs():
|
||||
model_detectors=[lambda hf_id: "krea-2" in hf_id.lower()],
|
||||
)
|
||||
# Qwen-Image
|
||||
register_configs(
|
||||
sampling_param_cls=QwenImage21SamplingParams,
|
||||
pipeline_config_cls=QwenImage21PipelineConfig,
|
||||
hf_model_paths=["Qwen/Qwen-Image-2.1"],
|
||||
model_detectors=[lambda hf_id: "qwen-image-2.1" in hf_id.lower()],
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=QwenImageSamplingParams,
|
||||
pipeline_config_cls=QwenImagePipelineConfig,
|
||||
@@ -1141,6 +1151,7 @@ def _register_configs():
|
||||
and "edit" not in hf_id.lower()
|
||||
and "layered" not in hf_id.lower()
|
||||
and "2512" not in hf_id.lower()
|
||||
and "qwen-image-2.1" not in hf_id.lower()
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
@@ -403,6 +403,10 @@ class CustomBlockAdapterSpec:
|
||||
|
||||
# Custom BlockAdapter metadata for models absent from cache-dit's registry.
|
||||
_CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, CustomBlockAdapterSpec] = {
|
||||
"QwenImage21Transformer2DModel": CustomBlockAdapterSpec(
|
||||
blocks_attr="transformer_blocks",
|
||||
forward_pattern=ForwardPattern.Pattern_3,
|
||||
),
|
||||
"ErnieImageTransformer2DModel": CustomBlockAdapterSpec(
|
||||
blocks_attr="layers",
|
||||
forward_pattern=ForwardPattern.Pattern_3,
|
||||
@@ -523,23 +527,19 @@ def enable_cache_on_transformer(
|
||||
"Please provide it in CacheDitConfig."
|
||||
)
|
||||
|
||||
# Prefer the standard path (transformer pre-registered in cache-dit). For
|
||||
# models absent from the registry, fall back to a manual BlockAdapter (see
|
||||
# _build_custom_block_adapter).
|
||||
custom_adapter = None
|
||||
if not BlockAdapterRegister.is_supported(transformer):
|
||||
custom_adapter = _build_custom_block_adapter(
|
||||
transformer, has_separate_cfg=has_separate_cfg
|
||||
# Native forward contracts take precedence over cache-dit's family-name matching.
|
||||
custom_adapter = _build_custom_block_adapter(
|
||||
transformer, has_separate_cfg=has_separate_cfg
|
||||
)
|
||||
if custom_adapter is None and not BlockAdapterRegister.is_supported(transformer):
|
||||
transformer_cls_name = transformer.__class__.__name__
|
||||
raise ValueError(
|
||||
f"{transformer_cls_name} is not officially supported by cache-dit. "
|
||||
"Supported cache-dit DiT families include Flux, QwenImage, HunyuanDiT, "
|
||||
"HunyuanVideo, Wan, CogVideoX, Mochi, and others. "
|
||||
"Please ensure your transformer belongs to one of these families or "
|
||||
"define a custom BlockAdapter."
|
||||
)
|
||||
if custom_adapter is None:
|
||||
transformer_cls_name = transformer.__class__.__name__
|
||||
raise ValueError(
|
||||
f"{transformer_cls_name} is not officially supported by cache-dit. "
|
||||
"Supported cache-dit DiT families include Flux, QwenImage, HunyuanDiT, "
|
||||
"HunyuanVideo, Wan, CogVideoX, Mochi, and others. "
|
||||
"Please ensure your transformer belongs to one of these families or "
|
||||
"define a custom BlockAdapter."
|
||||
)
|
||||
|
||||
# Build cache config (including SCM fields if provided)
|
||||
cache_config = DBCacheConfig(
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Transfer request extras with nested tensors through the normal tensor codec."""
|
||||
|
||||
import json
|
||||
|
||||
import torch
|
||||
from torch.utils._pytree import (
|
||||
tree_flatten,
|
||||
tree_unflatten,
|
||||
treespec_dumps,
|
||||
treespec_loads,
|
||||
)
|
||||
|
||||
|
||||
def extract_extra_tensors(extra, tensor_fields, scalar_fields):
|
||||
for key, value in extra.items():
|
||||
if key.startswith("_"):
|
||||
continue
|
||||
leaves, spec = tree_flatten(value)
|
||||
indices = [i for i, leaf in enumerate(leaves) if isinstance(leaf, torch.Tensor)]
|
||||
if not indices:
|
||||
try:
|
||||
json.dumps(value)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
continue
|
||||
scalar_fields[f"_extra_{key}"] = value
|
||||
continue
|
||||
tensors = [leaves[i] for i in indices]
|
||||
for i in indices:
|
||||
leaves[i] = None
|
||||
try:
|
||||
metadata = dict(spec=treespec_dumps(spec), leaves=leaves, indices=indices)
|
||||
json.dumps(metadata)
|
||||
except (TypeError, ValueError, OverflowError, NotImplementedError):
|
||||
continue
|
||||
name = f"_extra_tensor_tree_{key}"
|
||||
tensor_fields[name] = tensors
|
||||
scalar_fields[name] = metadata
|
||||
|
||||
|
||||
def restore_extra_tensors(extra, tensor_fields, scalar_fields):
|
||||
for name in list(scalar_fields):
|
||||
if not name.startswith("_extra_tensor_tree_"):
|
||||
continue
|
||||
metadata = scalar_fields.pop(name)
|
||||
leaves = metadata["leaves"]
|
||||
for index, tensor in zip(
|
||||
metadata["indices"], tensor_fields.pop(name), strict=True
|
||||
):
|
||||
leaves[index] = tensor
|
||||
extra[name[len("_extra_tensor_tree_") :]] = tree_unflatten(
|
||||
leaves, treespec_loads(metadata["spec"])
|
||||
)
|
||||
@@ -24,6 +24,10 @@ import torch
|
||||
import zmq
|
||||
|
||||
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||
from sglang.multimodal_gen.runtime.disaggregation.extra_tensors import (
|
||||
extract_extra_tensors,
|
||||
restore_extra_tensors,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.disaggregation.transport.buffer import (
|
||||
TransferTensorBuffer,
|
||||
@@ -205,18 +209,6 @@ def _is_default(value, field_info) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _extract_extra_fields(extra: dict, scalar_fields: dict) -> None:
|
||||
"""Extract JSON-serializable entries from Req.extra into scalar_fields."""
|
||||
for key, value in extra.items():
|
||||
if key.startswith("_"):
|
||||
continue
|
||||
try:
|
||||
json.dumps(value)
|
||||
scalar_fields[f"_extra_{key}"] = value
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
pass
|
||||
|
||||
|
||||
def _init_request_scheduler(scheduler: Any, req: Req, device: torch.device) -> None:
|
||||
extra_kwargs = {}
|
||||
mu = req.extra.get("mu") if hasattr(req, "extra") else None
|
||||
@@ -300,7 +292,7 @@ def extract_transfer_fields(req) -> tuple[dict, dict]:
|
||||
|
||||
extra = getattr(req, "extra", None)
|
||||
if extra:
|
||||
_extract_extra_fields(extra, scalar_fields)
|
||||
extract_extra_tensors(extra, tensor_fields, scalar_fields)
|
||||
|
||||
sp = getattr(req, "sampling_params", None)
|
||||
if sp is not None:
|
||||
@@ -1411,6 +1403,7 @@ class SchedulerDisaggMixin:
|
||||
object.__setattr__(req, f.name, f.default_factory())
|
||||
# Ensure sampling_params is not None so __getattr__ delegation works
|
||||
object.__setattr__(req, "sampling_params", SamplingParams())
|
||||
restore_extra_tensors(req.extra, tensors, scalar_fields)
|
||||
# Restore _extra_* prefixed fields into req.extra dict
|
||||
extra_keys = [k for k in scalar_fields if k.startswith("_extra_")]
|
||||
for key in extra_keys:
|
||||
|
||||
@@ -26,6 +26,8 @@ _DTYPE_TO_STR = {
|
||||
torch.int64: "int64",
|
||||
torch.uint8: "uint8",
|
||||
torch.bool: "bool",
|
||||
torch.complex64: "complex64",
|
||||
torch.complex128: "complex128",
|
||||
}
|
||||
_STR_TO_DTYPE = {v: k for k, v in _DTYPE_TO_STR.items()}
|
||||
|
||||
@@ -48,7 +50,7 @@ class TensorWrapper:
|
||||
"""Expose a CPU-contiguous tensor's data buffer for zero-copy ZMQ send."""
|
||||
|
||||
def __init__(self, tensor: torch.Tensor):
|
||||
if tensor.is_cuda or tensor.is_npu:
|
||||
if tensor.device.type != "cpu":
|
||||
tensor = tensor.cpu()
|
||||
if not tensor.is_contiguous():
|
||||
tensor = tensor.contiguous()
|
||||
@@ -186,7 +188,11 @@ def unpack_tensors(
|
||||
buf = frame.buffer if hasattr(frame, "buffer") else bytes(frame)
|
||||
dtype = str_to_dtype(desc.dtype)
|
||||
# clone() to own the memory (decouple from ZMQ buffer lifetime)
|
||||
tensor = torch.frombuffer(buf, dtype=dtype).reshape(desc.shape).clone()
|
||||
tensor = (
|
||||
torch.empty(desc.shape, dtype=dtype)
|
||||
if 0 in desc.shape
|
||||
else torch.frombuffer(buf, dtype=dtype).reshape(desc.shape).clone()
|
||||
)
|
||||
if device != "cpu" and device != torch.device("cpu"):
|
||||
tensor = tensor.to(device)
|
||||
|
||||
|
||||
@@ -0,0 +1,546 @@
|
||||
# Copyright 2026 Qwen-Image Team and The HuggingFace Team
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.kernels.ops.diffusion import (
|
||||
BitExactFusionGate,
|
||||
can_use_fused_complex_rope,
|
||||
can_use_fused_layernorm_modulate,
|
||||
can_use_fused_silu_mul,
|
||||
can_use_rmsnorm_preserve_reduction,
|
||||
fused_complex_rope,
|
||||
fused_layernorm_modulate,
|
||||
fused_silu_mul_bitexact,
|
||||
residual_gate_add,
|
||||
rmsnorm_preserve_reduction,
|
||||
tensors_equal,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.rope.qknorm_complex_rope_kv_triton import (
|
||||
can_use_qknorm_complex_rope_kv,
|
||||
qknorm_complex_rope_kv,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.rope.qknorm_complex_rope_triton import (
|
||||
can_use_qknorm_complex_rope,
|
||||
qknorm_complex_rope,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_sp_world_size,
|
||||
get_tp_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
get_sp_parallel_rank,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.attention import LocalAttention, USPAttention
|
||||
from sglang.multimodal_gen.runtime.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||
LayerwiseOffloadableModuleMixin,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||
from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_ROPE_FUSION = BitExactFusionGate("Qwen-Image 2.1 complex RoPE")
|
||||
_SILU_MUL_FUSION = BitExactFusionGate("Qwen-Image 2.1 SiLU-mul")
|
||||
_QK_ROPE_FUSION = BitExactFusionGate("Qwen-Image 2.1 Q/K RMSNorm + complex RoPE")
|
||||
_KV_ROPE_FUSION = BitExactFusionGate("Qwen-Image 2.1 K RMSNorm + RoPE + KV packing")
|
||||
_QK_NORM_FUSION = BitExactFusionGate("Qwen-Image 2.1 Q/K RMSNorm")
|
||||
_MODULATION_FUSION = BitExactFusionGate("Qwen-Image 2.1 LayerNorm modulation")
|
||||
|
||||
|
||||
def build_layout(image_slots, image_shapes, axes_dims, device):
|
||||
"""Expand each condition-image slot to its complete latent grid before denoising."""
|
||||
indices, image_indices, positions, segments = [], [], [], []
|
||||
cursor = position = image_index = 0
|
||||
for text_index, is_image in enumerate(image_slots):
|
||||
if not is_image:
|
||||
indices.append(text_index)
|
||||
positions.append((position, position, position))
|
||||
position += 1
|
||||
continue
|
||||
start = len(indices)
|
||||
if start > cursor:
|
||||
segments.append((cursor, start, False))
|
||||
_, height, width = image_shapes[image_index]
|
||||
for h in range(-(height - height // 2), height // 2):
|
||||
for w in range(-(width - width // 2), width // 2):
|
||||
indices.append(text_index)
|
||||
image_indices.append(len(indices) - 1)
|
||||
positions.append((position, h, w))
|
||||
segments.append((start, len(indices), True))
|
||||
cursor = len(indices)
|
||||
position += max(height, width)
|
||||
image_index += 1
|
||||
if image_index != len(image_shapes) - 1:
|
||||
raise ValueError("condition-image slots do not match image_shapes")
|
||||
if len(indices) > cursor:
|
||||
segments.append((cursor, len(indices), False))
|
||||
prefix_len = len(indices)
|
||||
_, height, width = image_shapes[-1]
|
||||
for h in range(-(height - height // 2), height // 2):
|
||||
for w in range(-(width - width // 2), width // 2):
|
||||
positions.append((position, h, w))
|
||||
pos = torch.tensor(positions, device=device, dtype=torch.float32)
|
||||
angles = torch.cat(
|
||||
[
|
||||
pos[:, axis : axis + 1]
|
||||
* (10000.0 ** (-torch.arange(0, dim, 2, device=device).float() / dim))
|
||||
for axis, dim in enumerate(axes_dims)
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
rope = torch.polar(torch.ones_like(angles), angles)
|
||||
return dict(
|
||||
text_indices=torch.tensor(indices, device=device, dtype=torch.long),
|
||||
image_indices=torch.tensor(image_indices, device=device, dtype=torch.long),
|
||||
prefix_rope=rope[:prefix_len],
|
||||
target_rope=rope[prefix_len:],
|
||||
segments=tuple(segments),
|
||||
)
|
||||
|
||||
|
||||
def apply_rope(x, rope):
|
||||
fused = None
|
||||
if can_use_fused_complex_rope(x, rope) and _ROPE_FUSION.can_attempt_once():
|
||||
fused = fused_complex_rope(x, rope)
|
||||
if _ROPE_FUSION.verified:
|
||||
return fused
|
||||
z = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
||||
out = torch.view_as_real(z * rope[None, :, None]).flatten(-2).to(x.dtype)
|
||||
if fused is not None:
|
||||
return _ROPE_FUSION.accept_or_fallback(fused, out, logger=logger)
|
||||
return out
|
||||
|
||||
|
||||
def apply_qk_norm(x, norm):
|
||||
fused = None
|
||||
if (
|
||||
can_use_rmsnorm_preserve_reduction(x, norm.weight)
|
||||
and _QK_NORM_FUSION.can_attempt_once()
|
||||
):
|
||||
fused = rmsnorm_preserve_reduction(x, norm.weight, norm.variance_epsilon)
|
||||
if _QK_NORM_FUSION.verified:
|
||||
return fused
|
||||
out = norm(x)
|
||||
if fused is not None:
|
||||
return _QK_NORM_FUSION.accept_or_fallback(fused, out, logger=logger)
|
||||
return out
|
||||
|
||||
|
||||
def apply_qk_norm_rope(x, norm, rope):
|
||||
fused = None
|
||||
if (
|
||||
can_use_qknorm_complex_rope(x, norm.weight, rope)
|
||||
and _QK_ROPE_FUSION.can_attempt_once()
|
||||
):
|
||||
fused = qknorm_complex_rope(x, norm.weight, rope, norm.variance_epsilon)
|
||||
if _QK_ROPE_FUSION.verified:
|
||||
return fused
|
||||
out = apply_rope(apply_qk_norm(x, norm), rope)
|
||||
if fused is not None:
|
||||
return _QK_ROPE_FUSION.accept_or_fallback(fused, out, logger=logger)
|
||||
return out
|
||||
|
||||
|
||||
def apply_modulation(x, norm, scale):
|
||||
fused = None
|
||||
if (
|
||||
can_use_fused_layernorm_modulate(x, scale.squeeze(1), None)
|
||||
and _MODULATION_FUSION.can_attempt_once()
|
||||
):
|
||||
fused = fused_layernorm_modulate(x, scale.squeeze(1), None, norm.eps)
|
||||
if _MODULATION_FUSION.verified:
|
||||
return fused
|
||||
out = norm(x) * (1 + scale)
|
||||
if fused is not None:
|
||||
return _MODULATION_FUSION.accept_or_fallback(fused, out, logger=logger)
|
||||
return out
|
||||
|
||||
|
||||
class QwenImage21ZeroCenterRMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.zeros(dim))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x):
|
||||
scale = self.weight.float() + 1
|
||||
value = x.float()
|
||||
return (
|
||||
value
|
||||
* torch.rsqrt(value.square().mean(-1, keepdim=True) + self.eps)
|
||||
* scale
|
||||
).to(x.dtype)
|
||||
|
||||
|
||||
class QwenImage21TextProjection(nn.Module):
|
||||
def __init__(self, context_dim, dim, eps):
|
||||
super().__init__()
|
||||
self.text_norm = QwenImage21ZeroCenterRMSNorm(context_dim, eps)
|
||||
self.in_layer = nn.Linear(context_dim, dim, bias=False)
|
||||
self.out_layer = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
return self.out_layer(
|
||||
nn.functional.gelu(self.in_layer(self.text_norm(x)), approximate="tanh")
|
||||
)
|
||||
|
||||
|
||||
class QwenImage21TimeEmbedding(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.timestep_embedder = nn.Module()
|
||||
self.timestep_embedder.linear_1 = nn.Linear(256, dim, bias=False)
|
||||
self.timestep_embedder.linear_2 = nn.Linear(dim, dim, bias=False)
|
||||
|
||||
def forward(self, t, dtype):
|
||||
freq = torch.exp(
|
||||
-math.log(10000) * torch.arange(128, device=t.device).float() / 128
|
||||
)
|
||||
angles = t.float()[:, None] * 1000 * freq
|
||||
x = torch.cat([angles.cos(), angles.sin()], dim=-1).to(dtype)
|
||||
return self.timestep_embedder.linear_2(
|
||||
nn.functional.silu(self.timestep_embedder.linear_1(x))
|
||||
)
|
||||
|
||||
|
||||
class QwenImage21FeedForward(nn.Module):
|
||||
def __init__(self, dim, ratio, quant_config, prefix):
|
||||
super().__init__()
|
||||
self.proj = ColumnParallelLinear(
|
||||
dim,
|
||||
dim * ratio,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.proj",
|
||||
)
|
||||
self.gate_layer = ColumnParallelLinear(
|
||||
dim,
|
||||
dim * ratio,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.gate_layer",
|
||||
)
|
||||
self.out = RowParallelLinear(
|
||||
dim * ratio,
|
||||
dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.out",
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
gate, value = self.gate_layer(x)[0], self.proj(x)[0]
|
||||
fused = None
|
||||
if can_use_fused_silu_mul(gate, value) and _SILU_MUL_FUSION.can_attempt_once():
|
||||
fused = fused_silu_mul_bitexact(gate, value)
|
||||
if _SILU_MUL_FUSION.verified:
|
||||
return self.out(fused)[0]
|
||||
hidden = nn.functional.silu(gate) * value
|
||||
if fused is not None:
|
||||
hidden = _SILU_MUL_FUSION.accept_or_fallback(fused, hidden, logger=logger)
|
||||
return self.out(hidden)[0]
|
||||
|
||||
|
||||
class QwenImage21Attention(nn.Module):
|
||||
def __init__(self, ac, quant_config, prefix):
|
||||
super().__init__()
|
||||
dim = ac.hidden_size
|
||||
self.heads = ac.num_attention_heads // get_tp_world_size()
|
||||
self.head_dim = ac.attention_head_dim
|
||||
self.to_q = ColumnParallelLinear(
|
||||
dim, dim, bias=False, quant_config=quant_config, prefix=f"{prefix}.to_q"
|
||||
)
|
||||
self.to_k = ColumnParallelLinear(
|
||||
dim, dim, bias=False, quant_config=quant_config, prefix=f"{prefix}.to_k"
|
||||
)
|
||||
self.to_v = ColumnParallelLinear(
|
||||
dim, dim, bias=False, quant_config=quant_config, prefix=f"{prefix}.to_v"
|
||||
)
|
||||
self.to_out = nn.ModuleList(
|
||||
[
|
||||
RowParallelLinear(
|
||||
dim,
|
||||
dim,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.to_out.0",
|
||||
)
|
||||
]
|
||||
)
|
||||
self.norm_q = RMSNorm(
|
||||
self.head_dim, ac.eps, cast_x_before_out_mul=True, force_native=True
|
||||
)
|
||||
self.norm_k = RMSNorm(
|
||||
self.head_dim, ac.eps, cast_x_before_out_mul=True, force_native=True
|
||||
)
|
||||
backends = QwenImage21Transformer2DModel._supported_attention_backends
|
||||
self.local_attn = LocalAttention(
|
||||
self.heads, self.head_dim, supported_attention_backends=backends
|
||||
)
|
||||
self.target_attn = USPAttention(
|
||||
self.heads, self.head_dim, supported_attention_backends=backends
|
||||
)
|
||||
|
||||
def project_qkv(self, x):
|
||||
q = self.to_q(x)[0].unflatten(-1, (self.heads, self.head_dim))
|
||||
k = self.to_k(x)[0].unflatten(-1, (self.heads, self.head_dim))
|
||||
v = self.to_v(x)[0].unflatten(-1, (self.heads, self.head_dim))
|
||||
return q, k, v
|
||||
|
||||
def qkv(self, x, rope):
|
||||
q, k, v = self.project_qkv(x)
|
||||
return (
|
||||
apply_qk_norm_rope(q, self.norm_q, rope),
|
||||
apply_qk_norm_rope(k, self.norm_k, rope),
|
||||
v,
|
||||
)
|
||||
|
||||
def forward(self, x, rope, prefix, prefix_rope, segments, cache):
|
||||
if cache:
|
||||
kp, vp = cache["key"], cache["value"]
|
||||
prefix_output = None
|
||||
else:
|
||||
qp, kp, vp = self.qkv(prefix, prefix_rope)
|
||||
outputs = []
|
||||
# text runs are causal; image blocks see the entire preceding sequence and themselves
|
||||
for start, end, is_image in segments:
|
||||
mask = None
|
||||
if not is_image:
|
||||
mask = (
|
||||
torch.arange(end, device=x.device)[None, :]
|
||||
<= torch.arange(start, end, device=x.device)[:, None]
|
||||
)
|
||||
mask = mask[None, None]
|
||||
outputs.append(
|
||||
self.local_attn(
|
||||
qp[:, start:end], kp[:, :end], vp[:, :end], attn_mask=mask
|
||||
)
|
||||
)
|
||||
prefix_output = self.to_out[0](torch.cat(outputs, dim=1).flatten(2))[0]
|
||||
if cache is not None:
|
||||
cache.update(key=kp, value=vp)
|
||||
q, k, v = self.project_qkv(x)
|
||||
q = apply_qk_norm_rope(q, self.norm_q, rope)
|
||||
packed = None
|
||||
if (
|
||||
get_sp_world_size() == 1
|
||||
and can_use_qknorm_complex_rope_kv(k, self.norm_k.weight, rope, v, kp, vp)
|
||||
and _KV_ROPE_FUSION.can_attempt_once()
|
||||
):
|
||||
packed = qknorm_complex_rope_kv(
|
||||
k, self.norm_k.weight, rope, v, kp, vp, self.norm_k.variance_epsilon
|
||||
)
|
||||
if not _KV_ROPE_FUSION.verified:
|
||||
reference = (
|
||||
torch.cat([kp, apply_rope(apply_qk_norm(k, self.norm_k), rope)], 1),
|
||||
torch.cat([vp, v], 1),
|
||||
)
|
||||
packed = _KV_ROPE_FUSION.accept_or_fallback(
|
||||
packed,
|
||||
reference,
|
||||
equal=tensors_equal,
|
||||
logger=logger,
|
||||
)
|
||||
if packed is not None:
|
||||
out = self.target_attn(q, *packed)
|
||||
else:
|
||||
k = apply_qk_norm_rope(k, self.norm_k, rope)
|
||||
out = self.target_attn.forward_with_replicated_kv_prefix(q, kp, vp, k, v)
|
||||
return self.to_out[0](out.flatten(2))[0], prefix_output
|
||||
|
||||
|
||||
class QwenImage21TransformerBlock(nn.Module):
|
||||
def __init__(self, ac, quant_config, prefix):
|
||||
super().__init__()
|
||||
self.img_norm1 = nn.LayerNorm(
|
||||
ac.hidden_size, eps=ac.eps, elementwise_affine=False
|
||||
)
|
||||
self.img_norm2 = nn.LayerNorm(
|
||||
ac.hidden_size, eps=ac.eps, elementwise_affine=False
|
||||
)
|
||||
self.attn = QwenImage21Attention(ac, quant_config, f"{prefix}.attn")
|
||||
self.img_mlp = QwenImage21FeedForward(
|
||||
ac.hidden_size, ac.mlp_ratio, quant_config, f"{prefix}.img_mlp"
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
modulation,
|
||||
prefix_state,
|
||||
prefix_modulation,
|
||||
layout,
|
||||
rope,
|
||||
cache,
|
||||
):
|
||||
prefix = prefix_state.get("hidden_states")
|
||||
scale1, gate1, scale2, gate2 = modulation
|
||||
p = None
|
||||
if not cache:
|
||||
ps1, pg1, ps2, pg2 = prefix_modulation
|
||||
p = apply_modulation(prefix, self.img_norm1, ps1)
|
||||
attention, prefix_attention = self.attn(
|
||||
apply_modulation(hidden_states, self.img_norm1, scale1),
|
||||
rope,
|
||||
p,
|
||||
layout["prefix_rope"],
|
||||
layout["segments"],
|
||||
cache,
|
||||
)
|
||||
hidden_states = residual_gate_add(hidden_states, attention, gate1)
|
||||
hidden_states = residual_gate_add(
|
||||
hidden_states,
|
||||
self.img_mlp(apply_modulation(hidden_states, self.img_norm2, scale2)),
|
||||
gate2,
|
||||
)
|
||||
if prefix_attention is not None:
|
||||
prefix = residual_gate_add(prefix, prefix_attention, pg1)
|
||||
prefix = residual_gate_add(
|
||||
prefix,
|
||||
self.img_mlp(apply_modulation(prefix, self.img_norm2, ps2)),
|
||||
pg2,
|
||||
)
|
||||
prefix_state["hidden_states"] = prefix
|
||||
return hidden_states
|
||||
|
||||
|
||||
class QwenImage21OutputNorm(nn.Module):
|
||||
def __init__(self, dim, eps):
|
||||
super().__init__()
|
||||
self.linear = nn.Linear(dim, dim, bias=False)
|
||||
self.norm = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||
|
||||
def forward(self, x, temb):
|
||||
return apply_modulation(
|
||||
x, self.norm, self.linear(nn.functional.silu(temb))[:, None]
|
||||
)
|
||||
|
||||
|
||||
class QwenImage21Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
_supported_attention_backends = {
|
||||
AttentionBackendEnum.FA,
|
||||
AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.SAGE_ATTN_3,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
}
|
||||
_fsdp_shard_conditions = [
|
||||
lambda name, module: isinstance(module, QwenImage21TransformerBlock)
|
||||
]
|
||||
_compile_conditions = _fsdp_shard_conditions
|
||||
layer_names = ["transformer_blocks"]
|
||||
param_names_mapping = {}
|
||||
|
||||
def __init__(self, config, hf_config, quant_config=None, **kwargs):
|
||||
super().__init__(config, hf_config=hf_config, **kwargs)
|
||||
ac = self.config
|
||||
if ac.patch_size != 1 or not ac.causal_condition or not ac.causal_block:
|
||||
raise ValueError(
|
||||
"Qwen-Image 2.1 requires patch_size=1, causal_condition=True and causal_block=True"
|
||||
)
|
||||
self.hidden_size = ac.hidden_size
|
||||
self.num_attention_heads = ac.num_attention_heads
|
||||
self.num_channels_latents = ac.in_channels
|
||||
self.img_in = nn.Linear(ac.in_channels, ac.hidden_size, bias=False)
|
||||
self.txt_in = QwenImage21TextProjection(
|
||||
ac.context_in_dim, ac.hidden_size, ac.eps
|
||||
)
|
||||
self.time_text_embed = QwenImage21TimeEmbedding(ac.hidden_size)
|
||||
self.modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(ac.hidden_size, ac.hidden_size * 4, bias=False)
|
||||
)
|
||||
self.transformer_blocks = nn.ModuleList(
|
||||
[
|
||||
QwenImage21TransformerBlock(ac, quant_config, f"transformer_blocks.{i}")
|
||||
for i in range(ac.num_layers)
|
||||
]
|
||||
)
|
||||
self.norm_out = QwenImage21OutputNorm(ac.hidden_size, ac.eps)
|
||||
self.proj_out = nn.Linear(ac.hidden_size, ac.out_channels, bias=False)
|
||||
|
||||
def prepare_modulation(self, temb):
|
||||
# All blocks share these gates. Preserve the native tanh and its dtype,
|
||||
# but compute it once per timestep instead of once per block.
|
||||
scale1, gate1, scale2, gate2 = self.modulation(temb)[:, None].chunk(4, dim=-1)
|
||||
return scale1, gate1.tanh(), scale2, gate2.tanh()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
timestep,
|
||||
layouts,
|
||||
condition_latents=None,
|
||||
prefix_caches=None,
|
||||
**kwargs,
|
||||
):
|
||||
if isinstance(encoder_hidden_states, list):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
sp = get_sp_world_size()
|
||||
target_len = hidden_states.shape[1]
|
||||
if target_len % sp:
|
||||
raise ValueError(
|
||||
f"target token count {target_len} must be divisible by SP degree {sp}"
|
||||
)
|
||||
local_len = target_len // sp
|
||||
rank = get_sp_parallel_rank()
|
||||
start, end = rank * local_len, (rank + 1) * local_len
|
||||
images = self.img_in(hidden_states[:, start:end])
|
||||
temb = self.time_text_embed((timestep.to(images.dtype) / 1000), images.dtype)
|
||||
modulation = self.prepare_modulation(temb)
|
||||
prefix_modulation = None
|
||||
if prefix_caches is None or any(not cache[0] for cache in prefix_caches):
|
||||
zero_temb = self.time_text_embed(
|
||||
timestep.new_zeros(1).to(images.dtype), images.dtype
|
||||
)
|
||||
prefix_modulation = self.prepare_modulation(zero_temb)
|
||||
outputs = []
|
||||
for sample, layout in enumerate(layouts):
|
||||
caches = (
|
||||
prefix_caches[sample]
|
||||
if prefix_caches is not None
|
||||
else [None] * len(self.transformer_blocks)
|
||||
)
|
||||
prefix = None
|
||||
if not caches[0]:
|
||||
prefix = self.txt_in(
|
||||
encoder_hidden_states[sample : sample + 1]
|
||||
).index_select(1, layout["text_indices"])
|
||||
if condition_latents is not None:
|
||||
prefix[:, layout["image_indices"]] = self.img_in(
|
||||
condition_latents[sample : sample + 1]
|
||||
)
|
||||
prefix_state = {"hidden_states": prefix}
|
||||
x = images[sample : sample + 1]
|
||||
sample_modulation = tuple(
|
||||
value[sample : sample + 1] for value in modulation
|
||||
)
|
||||
for i, block in enumerate(self.transformer_blocks):
|
||||
x = block(
|
||||
x,
|
||||
sample_modulation,
|
||||
prefix_state,
|
||||
prefix_modulation,
|
||||
layout,
|
||||
layout["target_rope"][start:end],
|
||||
caches[i],
|
||||
)
|
||||
outputs.append(self.proj_out(self.norm_out(x, temb[sample : sample + 1])))
|
||||
output = torch.cat(outputs)
|
||||
if sp > 1:
|
||||
output = sequence_model_parallel_all_gather(output, dim=1)
|
||||
return output
|
||||
|
||||
|
||||
EntryClass = QwenImage21Transformer2DModel
|
||||
@@ -1191,7 +1191,11 @@ class Qwen3VLModel(nn.Module):
|
||||
|
||||
|
||||
class Qwen3VLForConditionalGeneration(TextEncoder):
|
||||
layer_names = [*TextEncoder.layer_names, "model.visual.blocks"]
|
||||
layer_names = [
|
||||
*TextEncoder.layer_names,
|
||||
"model.visual.blocks",
|
||||
"model.visual.deepstack_merger_list",
|
||||
]
|
||||
default_bitsandbytes_target_modules = [
|
||||
".gate_up_proj.",
|
||||
".down_proj.",
|
||||
@@ -1216,8 +1220,11 @@ class Qwen3VLForConditionalGeneration(TextEncoder):
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__(config)
|
||||
quant_config = config.quant_config
|
||||
config = config.arch_config
|
||||
self.model = Qwen3VLModel(config)
|
||||
self.model = Qwen3VLModel(
|
||||
config, quant_config=quant_config, use_tensor_parallel=True, prefix="model"
|
||||
)
|
||||
self.lm_head = nn.Linear(
|
||||
config.text_config.hidden_size, config.text_config.vocab_size, bias=False
|
||||
)
|
||||
|
||||
@@ -30,12 +30,29 @@ class Qwen3VLVisionOutput:
|
||||
|
||||
|
||||
class Qwen3VLVisionRotaryEmbedding(nn.Module):
|
||||
recompute_on_device_change = False
|
||||
|
||||
def __init__(self, dim: int, theta: float = 10000.0) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.theta = theta
|
||||
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
self._inv_freq_device = inv_freq.device
|
||||
|
||||
def forward(self, sequence_length: int) -> torch.Tensor:
|
||||
if (
|
||||
self.recompute_on_device_change
|
||||
and self.inv_freq.device != self._inv_freq_device
|
||||
):
|
||||
# match resident initialization: CPU and GPU pow round differently
|
||||
indices = torch.arange(
|
||||
0, self.dim, 2, dtype=torch.float32, device=self.inv_freq.device
|
||||
)
|
||||
self.inv_freq = (1.0 / (self.theta ** (indices / self.dim))).to(
|
||||
self.inv_freq.dtype
|
||||
)
|
||||
self._inv_freq_device = self.inv_freq.device
|
||||
positions = torch.arange(
|
||||
sequence_length,
|
||||
device=self.inv_freq.device,
|
||||
@@ -184,6 +201,8 @@ def _vision_cu_seqlens(grid_thw: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
class Qwen3VLVisionTransformer(nn.Module):
|
||||
fp32_position_interpolation = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Any,
|
||||
@@ -250,7 +269,14 @@ class Qwen3VLVisionTransformer(nn.Module):
|
||||
num_grid_per_side=self.num_grid_per_side,
|
||||
spatial_merge_size=self.spatial_merge_size,
|
||||
)
|
||||
return (self.pos_embed(indices) * weights[:, :, None]).sum(0)
|
||||
if self.fp32_position_interpolation:
|
||||
return (self.pos_embed(indices) * weights[:, :, None]).sum(0)
|
||||
# Transformers 4.57 rounds each corner and each addition in the weight dtype
|
||||
corners = (
|
||||
self.pos_embed(indices)
|
||||
* weights.to(self.pos_embed.weight.dtype)[:, :, None]
|
||||
)
|
||||
return corners[0] + corners[1] + corners[2] + corners[3]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,774 @@
|
||||
# Copyright 2026 Qwen Team and The HuggingFace Team
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.kernels.ops.diffusion import dup_up3d_add
|
||||
from sglang.kernels.ops.diffusion.norm.channel_rmsnorm_preserve_reduction import (
|
||||
can_use_channel_rmsnorm,
|
||||
channel_rmsnorm_preserve_reduction,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.sites.bitexact_gate import BitExactFusionGate
|
||||
from sglang.multimodal_gen.configs.models.vaes.qwenimage21 import QwenImage21VAEConfig
|
||||
from sglang.multimodal_gen.runtime.distributed import (
|
||||
get_decode_parallel_rank,
|
||||
get_decode_parallel_world_size,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.layers.parallel_conv import (
|
||||
SpatialParallelConv2d,
|
||||
chunk_height_by_sizes,
|
||||
disable_spatial_parallel_decode,
|
||||
gather_and_trim_height,
|
||||
gather_variable_height,
|
||||
split_height_for_parallel_decode,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.vaes.common import (
|
||||
ParallelTiledVAE,
|
||||
can_install_spatial_shard_parallel_decode,
|
||||
should_run_spatial_shard_parallel_decode,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
_CHANNEL_RMSNORM_FUSION = BitExactFusionGate("Qwen-Image 2.1 VAE channel RMSNorm")
|
||||
|
||||
|
||||
def get_activation(name):
|
||||
if name != "silu":
|
||||
raise ValueError(f"unsupported VAE activation: {name}")
|
||||
return nn.SiLU()
|
||||
|
||||
|
||||
class QwenImage21AvgDown3D(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, factor_t, factor_s=1):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.factor_t = factor_t
|
||||
self.factor_s = factor_s
|
||||
self.factor = self.factor_t * self.factor_s * self.factor_s
|
||||
assert in_channels * self.factor % out_channels == 0
|
||||
self.group_size = in_channels * self.factor // out_channels
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t
|
||||
pad = (0, 0, 0, 0, pad_t, 0)
|
||||
x = F.pad(x, pad)
|
||||
B, C, T, H, W = x.shape
|
||||
x = x.view(
|
||||
B,
|
||||
C,
|
||||
T // self.factor_t,
|
||||
self.factor_t,
|
||||
H // self.factor_s,
|
||||
self.factor_s,
|
||||
W // self.factor_s,
|
||||
self.factor_s,
|
||||
)
|
||||
x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()
|
||||
x = x.view(
|
||||
B,
|
||||
C * self.factor,
|
||||
T // self.factor_t,
|
||||
H // self.factor_s,
|
||||
W // self.factor_s,
|
||||
)
|
||||
x = x.view(
|
||||
B,
|
||||
self.out_channels,
|
||||
self.group_size,
|
||||
T // self.factor_t,
|
||||
H // self.factor_s,
|
||||
W // self.factor_s,
|
||||
)
|
||||
x = x.mean(dim=2)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21DupUp3D(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, factor_t, factor_s=1):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.factor_t = factor_t
|
||||
self.factor_s = factor_s
|
||||
self.factor = self.factor_t * self.factor_s * self.factor_s
|
||||
assert out_channels * self.factor % in_channels == 0
|
||||
self.repeats = out_channels * self.factor // in_channels
|
||||
|
||||
def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor:
|
||||
x = x.repeat_interleave(self.repeats, dim=1)
|
||||
x = x.view(
|
||||
x.size(0),
|
||||
self.out_channels,
|
||||
self.factor_t,
|
||||
self.factor_s,
|
||||
self.factor_s,
|
||||
x.size(2),
|
||||
x.size(3),
|
||||
x.size(4),
|
||||
)
|
||||
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()
|
||||
x = x.view(
|
||||
x.size(0),
|
||||
self.out_channels,
|
||||
x.size(2) * self.factor_t,
|
||||
x.size(4) * self.factor_s,
|
||||
x.size(6) * self.factor_s,
|
||||
)
|
||||
if first_chunk:
|
||||
x = x[:, :, self.factor_t - 1 :, :, :]
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21CausalConv3d(nn.Conv2d):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size: int | tuple[int | int | int],
|
||||
stride: int | tuple[int | int | int] = 1,
|
||||
padding: int | tuple[int | int | int] = 0,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
)
|
||||
self._padding = (
|
||||
self.padding[1],
|
||||
self.padding[1],
|
||||
self.padding[0],
|
||||
self.padding[0],
|
||||
)
|
||||
self.padding = (0, 0)
|
||||
|
||||
def forward(self, x, cache_x=None):
|
||||
padding = list(self._padding)
|
||||
assert cache_x is None
|
||||
x = x.squeeze(2)
|
||||
x = F.pad(x, padding)
|
||||
x = super().forward(x)
|
||||
x = x.unsqueeze(2)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21RMS_norm(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
channel_first: bool = True,
|
||||
images: bool = True,
|
||||
bias: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
|
||||
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
|
||||
self.channel_first = channel_first
|
||||
self.scale = dim**0.5
|
||||
self.gamma = nn.Parameter(torch.ones(shape))
|
||||
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
|
||||
|
||||
def forward(self, x):
|
||||
fused = None
|
||||
if (
|
||||
self.channel_first
|
||||
and isinstance(self.bias, (int, float))
|
||||
and self.bias == 0
|
||||
and can_use_channel_rmsnorm(x, self.gamma)
|
||||
and _CHANNEL_RMSNORM_FUSION.can_attempt_once()
|
||||
):
|
||||
fused = channel_rmsnorm_preserve_reduction(x, self.gamma, self.scale)
|
||||
if _CHANNEL_RMSNORM_FUSION.verified:
|
||||
return fused
|
||||
normalized = F.normalize(
|
||||
x if x.dtype == torch.float64 else x.float(),
|
||||
dim=1 if self.channel_first else -1,
|
||||
).to(x.dtype)
|
||||
out = normalized * self.scale * self.gamma + self.bias
|
||||
if fused is not None:
|
||||
return _CHANNEL_RMSNORM_FUSION.accept_or_fallback(fused, out, logger=logger)
|
||||
return out
|
||||
|
||||
|
||||
class QwenImage21Upsample(nn.Upsample):
|
||||
def forward(self, x):
|
||||
return super().forward(x.float()).type_as(x)
|
||||
|
||||
|
||||
class QwenImage21Resample(nn.Module):
|
||||
def __init__(self, dim: int, mode: str, upsample_out_dim: int = None) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.mode = mode
|
||||
if upsample_out_dim is None:
|
||||
upsample_out_dim = dim // 2
|
||||
if mode == "upsample2d":
|
||||
self.resample = nn.Sequential(
|
||||
QwenImage21Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
|
||||
nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
|
||||
)
|
||||
elif mode == "upsample3d":
|
||||
self.resample = nn.Sequential(
|
||||
QwenImage21Upsample(scale_factor=(2.0, 2.0), mode="nearest-exact"),
|
||||
nn.Conv2d(dim, upsample_out_dim, 3, padding=1),
|
||||
)
|
||||
self.time_conv = QwenImage21CausalConv3d(
|
||||
dim, dim * 2, (1, 1), padding=(0, 0)
|
||||
)
|
||||
elif mode == "downsample2d":
|
||||
self.resample = nn.Sequential(
|
||||
nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))
|
||||
)
|
||||
elif mode == "downsample3d":
|
||||
self.resample = nn.Sequential(
|
||||
nn.ZeroPad2d((0, 1, 0, 1)), nn.Conv2d(dim, dim, 3, stride=(2, 2))
|
||||
)
|
||||
self.time_conv = QwenImage21CausalConv3d(
|
||||
dim, dim, (1, 1), stride=(1, 1), padding=(0, 0)
|
||||
)
|
||||
else:
|
||||
self.resample = nn.Identity()
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None):
|
||||
b, c, t, h, w = x.size()
|
||||
t = x.shape[2]
|
||||
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
||||
x = self.resample(x)
|
||||
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21ResidualBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim: int,
|
||||
out_dim: int,
|
||||
dropout: float = 0.0,
|
||||
non_linearity: str = "silu",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
self.norm1 = QwenImage21RMS_norm(in_dim, images=False)
|
||||
self.conv1 = QwenImage21CausalConv3d(in_dim, out_dim, 3, padding=1)
|
||||
self.norm2 = QwenImage21RMS_norm(out_dim, images=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.conv2 = QwenImage21CausalConv3d(out_dim, out_dim, 3, padding=1)
|
||||
self.conv_shortcut = (
|
||||
QwenImage21CausalConv3d(in_dim, out_dim, 1)
|
||||
if in_dim != out_dim
|
||||
else nn.Identity()
|
||||
)
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None):
|
||||
h = self.conv_shortcut(x)
|
||||
x = self.norm1(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv1(x)
|
||||
x = self.norm2(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.dropout(x)
|
||||
x = self.conv2(x)
|
||||
return x + h
|
||||
|
||||
|
||||
class QwenImage21AttentionBlock(nn.Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.spatial_parallel = False
|
||||
self.norm = QwenImage21RMS_norm(dim)
|
||||
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
|
||||
self.proj = nn.Conv2d(dim, dim, 1)
|
||||
|
||||
def forward(self, x):
|
||||
if self.spatial_parallel:
|
||||
x, heights = gather_variable_height(x)
|
||||
identity = x
|
||||
batch_size, channels, time, height, width = x.size()
|
||||
x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * time, channels, height, width)
|
||||
x = self.norm(x)
|
||||
qkv = self.to_qkv(x)
|
||||
qkv = qkv.reshape(batch_size * time, 1, channels * 3, -1)
|
||||
qkv = qkv.permute(0, 1, 3, 2).contiguous()
|
||||
q, k, v = qkv.chunk(3, dim=-1)
|
||||
x = F.scaled_dot_product_attention(q, k, v)
|
||||
x = (
|
||||
x.squeeze(1)
|
||||
.permute(0, 2, 1)
|
||||
.reshape(batch_size * time, channels, height, width)
|
||||
)
|
||||
x = self.proj(x)
|
||||
x = x.view(batch_size, time, channels, height, width)
|
||||
x = x.permute(0, 2, 1, 3, 4)
|
||||
x = x + identity
|
||||
return chunk_height_by_sizes(x, heights) if self.spatial_parallel else x
|
||||
|
||||
|
||||
class QwenImage21MidBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
dropout: float = 0.0,
|
||||
non_linearity: str = "silu",
|
||||
num_layers: int = 1,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
resnets = [QwenImage21ResidualBlock(dim, dim, dropout, non_linearity)]
|
||||
attentions = []
|
||||
for _ in range(num_layers):
|
||||
attentions.append(QwenImage21AttentionBlock(dim))
|
||||
resnets.append(QwenImage21ResidualBlock(dim, dim, dropout, non_linearity))
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None):
|
||||
x = self.resnets[0](x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
x = attn(x)
|
||||
x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21ResidualDownBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim,
|
||||
out_dim,
|
||||
dropout,
|
||||
num_res_blocks,
|
||||
temperal_downsample=False,
|
||||
down_flag=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.avg_shortcut = QwenImage21AvgDown3D(
|
||||
in_dim,
|
||||
out_dim,
|
||||
factor_t=2 if temperal_downsample else 1,
|
||||
factor_s=2 if down_flag else 1,
|
||||
)
|
||||
resnets = []
|
||||
for _ in range(num_res_blocks):
|
||||
resnets.append(QwenImage21ResidualBlock(in_dim, out_dim, dropout))
|
||||
in_dim = out_dim
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
if down_flag:
|
||||
mode = "downsample3d" if temperal_downsample else "downsample2d"
|
||||
self.downsampler = QwenImage21Resample(out_dim, mode=mode)
|
||||
else:
|
||||
self.downsampler = None
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None):
|
||||
x_copy = x
|
||||
for resnet in self.resnets:
|
||||
x = resnet(x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
if self.downsampler is not None:
|
||||
x = self.downsampler(x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
return x + self.avg_shortcut(x_copy)
|
||||
|
||||
|
||||
class QwenImage21Encoder3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
dim=128,
|
||||
z_dim=4,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_downsample=[True, True, False],
|
||||
dropout=0.0,
|
||||
non_linearity: str = "silu",
|
||||
is_residual: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_downsample = temperal_downsample
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
dims = [dim * u for u in [1] + dim_mult]
|
||||
scale = 1.0
|
||||
self.conv_in = QwenImage21CausalConv3d(in_channels, dims[0], 3, padding=1)
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
if is_residual:
|
||||
self.down_blocks.append(
|
||||
QwenImage21ResidualDownBlock(
|
||||
in_dim,
|
||||
out_dim,
|
||||
dropout,
|
||||
num_res_blocks,
|
||||
temperal_downsample=(
|
||||
temperal_downsample[i] if i != len(dim_mult) - 1 else False
|
||||
),
|
||||
down_flag=i != len(dim_mult) - 1,
|
||||
)
|
||||
)
|
||||
else:
|
||||
for _ in range(num_res_blocks):
|
||||
self.down_blocks.append(
|
||||
QwenImage21ResidualBlock(in_dim, out_dim, dropout)
|
||||
)
|
||||
if scale in attn_scales:
|
||||
self.down_blocks.append(QwenImage21AttentionBlock(out_dim))
|
||||
in_dim = out_dim
|
||||
if i != len(dim_mult) - 1:
|
||||
mode = "downsample3d" if temperal_downsample[i] else "downsample2d"
|
||||
self.down_blocks.append(QwenImage21Resample(out_dim, mode=mode))
|
||||
scale /= 2.0
|
||||
self.mid_block = QwenImage21MidBlock(
|
||||
out_dim, dropout, non_linearity, num_layers=1
|
||||
)
|
||||
self.norm_out = QwenImage21RMS_norm(out_dim, images=False)
|
||||
self.conv_out = QwenImage21CausalConv3d(out_dim, z_dim, 3, padding=1)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None):
|
||||
x = self.conv_in(x)
|
||||
for layer in self.down_blocks:
|
||||
x = layer(x)
|
||||
x = self.mid_block(x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21ResidualUpBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim: int,
|
||||
out_dim: int,
|
||||
num_res_blocks: int,
|
||||
dropout: float = 0.0,
|
||||
temperal_upsample: bool = False,
|
||||
up_flag: bool = False,
|
||||
non_linearity: str = "silu",
|
||||
):
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
if up_flag:
|
||||
self.avg_shortcut = QwenImage21DupUp3D(
|
||||
in_dim, out_dim, factor_t=2 if temperal_upsample else 1, factor_s=2
|
||||
)
|
||||
else:
|
||||
self.avg_shortcut = None
|
||||
resnets = []
|
||||
current_dim = in_dim
|
||||
for _ in range(num_res_blocks + 1):
|
||||
resnets.append(
|
||||
QwenImage21ResidualBlock(current_dim, out_dim, dropout, non_linearity)
|
||||
)
|
||||
current_dim = out_dim
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
if up_flag:
|
||||
upsample_mode = "upsample3d" if temperal_upsample else "upsample2d"
|
||||
self.upsampler = QwenImage21Resample(
|
||||
out_dim, mode=upsample_mode, upsample_out_dim=out_dim
|
||||
)
|
||||
else:
|
||||
self.upsampler = None
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False):
|
||||
x_copy = x
|
||||
for resnet in self.resnets:
|
||||
x = resnet(x)
|
||||
if self.upsampler is not None:
|
||||
x = self.upsampler(x)
|
||||
if self.avg_shortcut is not None:
|
||||
shortcut = self.avg_shortcut
|
||||
if (
|
||||
type(shortcut) is QwenImage21DupUp3D
|
||||
and x.is_cuda
|
||||
and x.dtype in (torch.float16, torch.bfloat16, torch.float32)
|
||||
and not torch.compiler.is_compiling()
|
||||
):
|
||||
fused = dup_up3d_add(
|
||||
x,
|
||||
x_copy,
|
||||
shortcut.factor_t,
|
||||
shortcut.factor_s,
|
||||
shortcut.repeats,
|
||||
first_chunk,
|
||||
)
|
||||
if fused is not None:
|
||||
return fused
|
||||
x = x + self.avg_shortcut(x_copy, first_chunk=first_chunk)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21UpBlock(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
in_dim: int,
|
||||
out_dim: int,
|
||||
num_res_blocks: int,
|
||||
dropout: float = 0.0,
|
||||
upsample_mode: str | None = None,
|
||||
non_linearity: str = "silu",
|
||||
):
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
resnets = []
|
||||
current_dim = in_dim
|
||||
for _ in range(num_res_blocks + 1):
|
||||
resnets.append(
|
||||
QwenImage21ResidualBlock(current_dim, out_dim, dropout, non_linearity)
|
||||
)
|
||||
current_dim = out_dim
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
self.upsamplers = None
|
||||
if upsample_mode is not None:
|
||||
self.upsamplers = nn.ModuleList(
|
||||
[QwenImage21Resample(out_dim, mode=upsample_mode)]
|
||||
)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=None):
|
||||
for resnet in self.resnets:
|
||||
x = resnet(x)
|
||||
if self.upsamplers is not None:
|
||||
x = self.upsamplers[0](x)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21Decoder3d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim=128,
|
||||
z_dim=4,
|
||||
dim_mult=[1, 2, 4, 4],
|
||||
num_res_blocks=2,
|
||||
attn_scales=[],
|
||||
temperal_upsample=[False, True, True],
|
||||
dropout=0.0,
|
||||
non_linearity: str = "silu",
|
||||
out_channels: int = 3,
|
||||
is_residual: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
self.dim_mult = dim_mult
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_upsample = temperal_upsample
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
||||
self.conv_in = QwenImage21CausalConv3d(z_dim, dims[0], 3, padding=1)
|
||||
self.mid_block = QwenImage21MidBlock(
|
||||
dims[0], dropout, non_linearity, num_layers=1
|
||||
)
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
if i > 0 and (not is_residual):
|
||||
in_dim = in_dim // 2
|
||||
up_flag = i != len(dim_mult) - 1
|
||||
upsample_mode = None
|
||||
if up_flag and temperal_upsample[i]:
|
||||
upsample_mode = "upsample3d"
|
||||
elif up_flag:
|
||||
upsample_mode = "upsample2d"
|
||||
if is_residual:
|
||||
up_block = QwenImage21ResidualUpBlock(
|
||||
in_dim=in_dim,
|
||||
out_dim=out_dim,
|
||||
num_res_blocks=num_res_blocks,
|
||||
dropout=dropout,
|
||||
temperal_upsample=temperal_upsample[i] if up_flag else False,
|
||||
up_flag=up_flag,
|
||||
non_linearity=non_linearity,
|
||||
)
|
||||
else:
|
||||
up_block = QwenImage21UpBlock(
|
||||
in_dim=in_dim,
|
||||
out_dim=out_dim,
|
||||
num_res_blocks=num_res_blocks,
|
||||
dropout=dropout,
|
||||
upsample_mode=upsample_mode,
|
||||
non_linearity=non_linearity,
|
||||
)
|
||||
self.up_blocks.append(up_block)
|
||||
self.norm_out = QwenImage21RMS_norm(out_dim, images=False)
|
||||
self.conv_out = QwenImage21CausalConv3d(out_dim, out_channels, 3, padding=1)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=None, first_chunk=False):
|
||||
x = self.conv_in(x)
|
||||
x = self.mid_block(x, feat_cache=feat_cache, feat_idx=feat_idx)
|
||||
for up_block in self.up_blocks:
|
||||
x = up_block(
|
||||
x, feat_cache=feat_cache, feat_idx=feat_idx, first_chunk=first_chunk
|
||||
)
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
def _patchify(x, patch_size):
|
||||
if patch_size == 1:
|
||||
return x
|
||||
if x.dim() != 5:
|
||||
raise ValueError(f"Invalid input shape: {x.shape}")
|
||||
batch_size, channels, frames, height, width = x.shape
|
||||
if height % patch_size != 0 or width % patch_size != 0:
|
||||
raise ValueError(
|
||||
f"Height ({height}) and width ({width}) must be divisible by patch_size ({patch_size})"
|
||||
)
|
||||
x = x.view(
|
||||
batch_size,
|
||||
channels,
|
||||
frames,
|
||||
height // patch_size,
|
||||
patch_size,
|
||||
width // patch_size,
|
||||
patch_size,
|
||||
)
|
||||
x = x.permute(0, 1, 6, 4, 2, 3, 5).contiguous()
|
||||
x = x.view(
|
||||
batch_size,
|
||||
channels * patch_size * patch_size,
|
||||
frames,
|
||||
height // patch_size,
|
||||
width // patch_size,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
def _unpatchify(x, patch_size):
|
||||
if patch_size == 1:
|
||||
return x
|
||||
if x.dim() != 5:
|
||||
raise ValueError(f"Invalid input shape: {x.shape}")
|
||||
batch_size, c_patches, frames, height, width = x.shape
|
||||
channels = c_patches // (patch_size * patch_size)
|
||||
x = x.view(batch_size, channels, patch_size, patch_size, frames, height, width)
|
||||
x = x.permute(0, 1, 4, 5, 3, 6, 2).contiguous()
|
||||
x = x.view(batch_size, channels, frames, height * patch_size, width * patch_size)
|
||||
return x
|
||||
|
||||
|
||||
class QwenImage21SpatialConv3d(SpatialParallelConv2d):
|
||||
def forward(self, x, cache_x=None):
|
||||
assert cache_x is None
|
||||
return super().forward(x.squeeze(2)).unsqueeze(2)
|
||||
|
||||
|
||||
def enable_qwen21_spatial_decode(module):
|
||||
for name, child in list(module.named_children()):
|
||||
if isinstance(child, QwenImage21AttentionBlock):
|
||||
# attention needs the full image; its pointwise projections stay local
|
||||
child.spatial_parallel = True
|
||||
elif isinstance(child, nn.Conv2d):
|
||||
causal = isinstance(child, QwenImage21CausalConv3d)
|
||||
conv_cls = QwenImage21SpatialConv3d if causal else SpatialParallelConv2d
|
||||
padding = (
|
||||
(child._padding[2], child._padding[0]) if causal else child.padding
|
||||
)
|
||||
conv = conv_cls(
|
||||
child.in_channels,
|
||||
child.out_channels,
|
||||
child.kernel_size,
|
||||
stride=child.stride,
|
||||
padding=padding,
|
||||
dilation=child.dilation,
|
||||
groups=child.groups,
|
||||
bias=child.bias is not None,
|
||||
)
|
||||
conv.weight, conv.bias = child.weight, child.bias
|
||||
setattr(module, name, conv)
|
||||
else:
|
||||
enable_qwen21_spatial_decode(child)
|
||||
|
||||
|
||||
class AutoencoderKLQwenImage21(ParallelTiledVAE):
|
||||
layer_names = [
|
||||
*ParallelTiledVAE.layer_names,
|
||||
"encoder.mid_block.resnets",
|
||||
"encoder.mid_block.attentions",
|
||||
"decoder.mid_block.resnets",
|
||||
"decoder.mid_block.attentions",
|
||||
]
|
||||
|
||||
def __init__(self, config: QwenImage21VAEConfig, **kwargs):
|
||||
super().__init__(config, **kwargs)
|
||||
ac = config.arch_config
|
||||
shared = dict(
|
||||
z_dim=ac.z_dim,
|
||||
dim_mult=list(ac.dim_mult),
|
||||
num_res_blocks=ac.num_res_blocks,
|
||||
attn_scales=list(ac.attn_scales),
|
||||
dropout=ac.dropout,
|
||||
is_residual=ac.is_residual,
|
||||
)
|
||||
if config.load_encoder:
|
||||
self.encoder = QwenImage21Encoder3d(
|
||||
in_channels=ac.in_channels,
|
||||
dim=ac.base_dim,
|
||||
**dict(shared, z_dim=ac.z_dim * 2),
|
||||
temperal_downsample=list(ac.temperal_downsample),
|
||||
)
|
||||
self.quant_conv = QwenImage21CausalConv3d(ac.z_dim * 2, ac.z_dim * 2, 1)
|
||||
if config.load_decoder:
|
||||
self.post_quant_conv = QwenImage21CausalConv3d(ac.z_dim, ac.z_dim, 1)
|
||||
self.decoder = QwenImage21Decoder3d(
|
||||
dim=ac.decoder_base_dim or ac.base_dim,
|
||||
**shared,
|
||||
temperal_upsample=list(ac.temperal_downsample)[::-1],
|
||||
out_channels=ac.out_channels,
|
||||
)
|
||||
self.spatial_parallel = (
|
||||
config.load_decoder and can_install_spatial_shard_parallel_decode(config)
|
||||
)
|
||||
if self.spatial_parallel:
|
||||
enable_qwen21_spatial_decode(self.decoder)
|
||||
|
||||
def _encode(self, x):
|
||||
if x.shape[2] != 1:
|
||||
raise ValueError("Qwen-Image 2.1 VAE expects one image frame")
|
||||
if self.config.patch_size is not None:
|
||||
x = _patchify(x, self.config.patch_size)
|
||||
return self.quant_conv(self.encoder(x))
|
||||
|
||||
def _decode(self, z):
|
||||
if z.shape[2] != 1:
|
||||
raise ValueError("Qwen-Image 2.1 VAE expects one latent frame")
|
||||
z = self.post_quant_conv(z)
|
||||
parallel = self.spatial_parallel and should_run_spatial_shard_parallel_decode(
|
||||
self.config, z
|
||||
)
|
||||
if parallel:
|
||||
z, expected_height = split_height_for_parallel_decode(
|
||||
z,
|
||||
expected_height=z.shape[-2] * self.spatial_compression_ratio,
|
||||
world_size=get_decode_parallel_world_size(),
|
||||
rank=get_decode_parallel_rank(),
|
||||
)
|
||||
x = self.decoder(z, first_chunk=True)
|
||||
else:
|
||||
with disable_spatial_parallel_decode():
|
||||
x = self.decoder(z, first_chunk=True)
|
||||
if self.config.patch_size is not None:
|
||||
x = _unpatchify(x, self.config.patch_size)
|
||||
if parallel:
|
||||
x = gather_and_trim_height(x, expected_height)
|
||||
return x.clamp(-1, 1)
|
||||
|
||||
|
||||
EntryClass = AutoencoderKLQwenImage21
|
||||
@@ -0,0 +1,52 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image21 import (
|
||||
QwenImage21DenoisingStage,
|
||||
QwenImage21EncodingStage,
|
||||
QwenImage21InputValidationStage,
|
||||
prepare_qwen21_mu,
|
||||
)
|
||||
|
||||
|
||||
class QwenImage21Pipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
pipeline_name = "QwenImage21Pipeline"
|
||||
_required_config_modules = [
|
||||
"processor",
|
||||
"text_encoder",
|
||||
"transformer",
|
||||
"vae",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, server_args):
|
||||
self.add_stage(QwenImage21InputValidationStage())
|
||||
self.add_stage_factory(
|
||||
RoleType.ENCODER,
|
||||
lambda: QwenImage21EncodingStage(
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("processor"),
|
||||
self.get_module("vae"),
|
||||
self.get_module("scheduler"),
|
||||
),
|
||||
"conditioning_stage",
|
||||
)
|
||||
self.add_standard_latent_preparation_stage()
|
||||
self.add_standard_timestep_preparation_stage(
|
||||
prepare_extra_kwargs=[prepare_qwen21_mu]
|
||||
)
|
||||
self.add_stage_factory(
|
||||
RoleType.DENOISER,
|
||||
lambda: QwenImage21DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
"denoising_stage",
|
||||
)
|
||||
self.add_standard_decoding_stage()
|
||||
|
||||
|
||||
EntryClass = QwenImage21Pipeline
|
||||
@@ -263,6 +263,7 @@ class ComposedPipelineBase(ABC):
|
||||
"Flux2KleinPipeline": {"vae"},
|
||||
"QwenImageEditPipeline": {"vae"},
|
||||
"QwenImageEditPlusPipeline": {"vae"},
|
||||
"QwenImage21Pipeline": {"vae"},
|
||||
"QwenImageLayeredPipeline": {"vae", "transformer"},
|
||||
"LongCatImageEditPipeline": {"vae"},
|
||||
"GlmImagePipeline": {"vae", "transformer"},
|
||||
|
||||
@@ -77,6 +77,9 @@ class InputValidationStage(PipelineStage):
|
||||
super().__init__()
|
||||
self.vae_image_processor = vae_image_processor
|
||||
|
||||
def load_condition_image(self, image):
|
||||
return load_image(image)
|
||||
|
||||
def iter_sequential_requests(
|
||||
self, batch: Req, server_args: ServerArgs
|
||||
) -> Iterator[Req]:
|
||||
@@ -429,7 +432,7 @@ class InputValidationStage(PipelineStage):
|
||||
if path.endswith(".mp4"):
|
||||
image = load_video(path)[0]
|
||||
else:
|
||||
image = load_image(path)
|
||||
image = self.load_condition_image(path)
|
||||
batch.condition_image.append(image)
|
||||
|
||||
# Use the first image for size reference
|
||||
@@ -443,7 +446,7 @@ class InputValidationStage(PipelineStage):
|
||||
if batch.image_path.endswith(".mp4"):
|
||||
image = load_video(batch.image_path)[0]
|
||||
else:
|
||||
image = load_image(batch.image_path)
|
||||
image = self.load_condition_image(batch.image_path)
|
||||
batch.condition_image = image
|
||||
condition_image_width, condition_image_height = (
|
||||
image.width,
|
||||
|
||||
+259
@@ -0,0 +1,259 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import math
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||
ComponentUse,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image21 import build_layout
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils import (
|
||||
calculate_linear_shift,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
||||
InputValidationStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.vision import load_image
|
||||
|
||||
SYSTEM_PROMPT = "Comprehend and analyze the provided prompt."
|
||||
SYSTEM_TEMPLATE = f"<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n"
|
||||
|
||||
|
||||
def collapse_image_slots(hidden, input_ids, image_token_id):
|
||||
image_mask = input_ids == image_token_id
|
||||
keep = ~image_mask
|
||||
keep[0] = True
|
||||
keep[1:] |= image_mask[1:] & ~image_mask[:-1]
|
||||
return hidden[keep], image_mask[keep]
|
||||
|
||||
|
||||
class QwenImage21InputValidationStage(InputValidationStage):
|
||||
def load_condition_image(self, image):
|
||||
return load_image(image, convert_method=lambda image: image.convert("RGBA"))
|
||||
|
||||
def preprocess_condition_image(
|
||||
self, batch, server_args, condition_image_width, condition_image_height
|
||||
):
|
||||
# one model-owned resize is shared by the VLM and VAE in the encoding stage
|
||||
return None
|
||||
|
||||
def forward(self, batch, server_args):
|
||||
if batch.prompt is None:
|
||||
raise ValueError(
|
||||
"Qwen-Image 2.1 requires a prompt to build image-token positions"
|
||||
)
|
||||
batch = super().forward(batch, server_args)
|
||||
if batch.height % 32 or batch.width % 32:
|
||||
raise ValueError("Qwen-Image 2.1 height and width must be divisible by 32")
|
||||
return batch
|
||||
|
||||
|
||||
class QwenImage21EncodingStage(PipelineStage):
|
||||
def __init__(self, text_encoder, processor, vae, scheduler):
|
||||
super().__init__()
|
||||
self.text_encoder, self.processor, self.vae, self.scheduler = (
|
||||
text_encoder,
|
||||
processor,
|
||||
vae,
|
||||
scheduler,
|
||||
)
|
||||
self.text_encoder.model.visual.fp32_position_interpolation = False
|
||||
self.text_encoder.model.visual.rotary_pos_emb.recompute_on_device_change = True
|
||||
self.image_token_id = processor.tokenizer.convert_tokens_to_ids("<|image_pad|>")
|
||||
system_message = [
|
||||
{"role": "system", "content": [{"type": "text", "text": SYSTEM_PROMPT}]}
|
||||
]
|
||||
self.drop_idx = len(
|
||||
processor.apply_chat_template(
|
||||
system_message, tokenize=True, return_dict=False
|
||||
)[0]
|
||||
)
|
||||
|
||||
def component_uses(self, server_args, stage_name=None):
|
||||
name = self._component_stage_name(stage_name)
|
||||
return [
|
||||
# preserve the loader's mixed weight and rotary buffer dtypes
|
||||
ComponentUse(name, "text_encoder"),
|
||||
ComponentUse(name, "vae", target_dtype=torch.bfloat16),
|
||||
]
|
||||
|
||||
def encode_prompt(self, prompt, images, device):
|
||||
prefix = " ".join(
|
||||
f"<image{i + 1}><|vision_start|><|image_pad|><|vision_end|>"
|
||||
for i in range(len(images))
|
||||
)
|
||||
text = (
|
||||
SYSTEM_TEMPLATE
|
||||
+ f"<|im_start|>user\n{prefix}{prompt or ' '}<|im_end|>\n<|im_start|>assistant\n"
|
||||
)
|
||||
kwargs = dict(
|
||||
text=[text], padding=True, padding_side="left", return_tensors="pt"
|
||||
)
|
||||
if images:
|
||||
vision_images = []
|
||||
for image in images:
|
||||
if image.mode == "RGBA":
|
||||
# vision conditioning uses white compositing; the VAE keeps RGBA
|
||||
white = Image.new("RGB", image.size, (255, 255, 255))
|
||||
white.paste(image, mask=image.getchannel("A"))
|
||||
image = white
|
||||
vision_images.append(image)
|
||||
kwargs["images"] = vision_images
|
||||
inputs = self.processor(**kwargs).to(device)
|
||||
with self.use_declared_component(
|
||||
component_name="text_encoder", module=self.text_encoder
|
||||
) as encoder:
|
||||
outputs = encoder(
|
||||
**inputs, output_hidden_states=True, use_cache=False, logits_to_keep=1
|
||||
)
|
||||
# the checkpoint expects Transformers 4.57's pre-final-norm hidden state
|
||||
final_hidden = outputs.hidden_states[-1]
|
||||
valid = inputs.attention_mask[0].bool()
|
||||
hidden = final_hidden[0, valid][self.drop_idx :]
|
||||
ids = inputs.input_ids[0, valid][self.drop_idx :]
|
||||
return collapse_image_slots(hidden, ids, self.image_token_id)
|
||||
|
||||
def forward(self, batch, server_args):
|
||||
config = server_args.pipeline_config
|
||||
ac = config.vae_config.arch_config
|
||||
device = get_local_torch_device()
|
||||
images = batch.condition_image
|
||||
images = (
|
||||
[] if images is None else images if isinstance(images, list) else [images]
|
||||
)
|
||||
resized, shapes, conditions = [], [], []
|
||||
area = batch.height * batch.width
|
||||
image_mode = "RGBA" if ac.in_channels == 4 else "RGB"
|
||||
for image in images:
|
||||
if not isinstance(image, Image.Image):
|
||||
image = load_image(
|
||||
image, convert_method=lambda image: image.convert(image_mode)
|
||||
)
|
||||
width = max(
|
||||
32, round(math.sqrt(area * image.width / image.height) / 32) * 32
|
||||
)
|
||||
height = max(
|
||||
32, round(math.sqrt(area * image.height / image.width) / 32) * 32
|
||||
)
|
||||
resized.append(
|
||||
image.convert(image_mode).resize(
|
||||
(width, height), Image.Resampling.LANCZOS
|
||||
)
|
||||
)
|
||||
shapes.append((1, height // 16, width // 16))
|
||||
if resized:
|
||||
with self.use_declared_component(
|
||||
component_name="vae", module=self.vae
|
||||
) as vae:
|
||||
vae.use_tiling = config.vae_tiling
|
||||
for image in resized:
|
||||
pixels = torch.frombuffer(
|
||||
bytearray(image.tobytes()), dtype=torch.uint8
|
||||
).reshape(image.height, image.width, ac.in_channels)
|
||||
# preserve the reference's batch stride for identical cuDNN convolution rounding
|
||||
pixels = (
|
||||
pixels[None].permute(0, 3, 1, 2).unsqueeze(2).float() / 255.0
|
||||
)
|
||||
pixels = (2 * pixels - 1).to(device=device, dtype=torch.bfloat16)
|
||||
latent = vae.encode(pixels).mode()
|
||||
mean = latent.new_tensor(ac.latents_mean).view(1, ac.z_dim, 1, 1, 1)
|
||||
std = latent.new_tensor(ac.latents_std).view(1, ac.z_dim, 1, 1, 1)
|
||||
conditions.append(
|
||||
((latent - mean) / std).flatten(2).transpose(1, 2)
|
||||
)
|
||||
shapes.append((1, batch.height // 16, batch.width // 16))
|
||||
prompts = batch.prompt if isinstance(batch.prompt, list) else [batch.prompt]
|
||||
negatives = (
|
||||
batch.negative_prompt
|
||||
if isinstance(batch.negative_prompt, list)
|
||||
else [batch.negative_prompt] * len(prompts)
|
||||
)
|
||||
sample_count = len(prompts) * batch.num_outputs_per_prompt
|
||||
condition_latents = (
|
||||
torch.cat(conditions, dim=1).expand(sample_count, -1, -1)
|
||||
if conditions
|
||||
else None
|
||||
)
|
||||
for negative in [False, True] if batch.do_classifier_free_guidance else [False]:
|
||||
embeds, masks, layouts = [], [], []
|
||||
for prompt in negatives if negative else prompts:
|
||||
with set_forward_context(
|
||||
current_timestep=None, attn_metadata=None, forward_batch=batch
|
||||
):
|
||||
hidden, slots = self.encode_prompt(prompt, resized, device)
|
||||
layout = build_layout(
|
||||
slots.tolist(), shapes, config.dit_config.axes_dims_rope, device
|
||||
)
|
||||
for _ in range(batch.num_outputs_per_prompt):
|
||||
embeds.append(hidden)
|
||||
layouts.append(layout)
|
||||
max_length = max(x.shape[0] for x in embeds)
|
||||
for x in embeds:
|
||||
masks.append(torch.arange(max_length, device=device) < x.shape[0])
|
||||
packed = torch.stack(
|
||||
[
|
||||
torch.nn.functional.pad(x, (0, 0, 0, max_length - x.shape[0]))
|
||||
for x in embeds
|
||||
]
|
||||
)
|
||||
mask = torch.stack(masks)
|
||||
if negative:
|
||||
batch.negative_prompt_embeds = [packed]
|
||||
batch.negative_prompt_embeds_mask = [mask]
|
||||
batch.negative_prompt_seq_lens = [mask.sum(1).tolist()]
|
||||
else:
|
||||
batch.prompt_embeds = [packed]
|
||||
batch.prompt_embeds_mask = [mask]
|
||||
batch.prompt_seq_lens = [mask.sum(1).tolist()]
|
||||
batch.extra["qwen21_negative" if negative else "qwen21_positive"] = dict(
|
||||
layouts=layouts,
|
||||
condition_latents=condition_latents,
|
||||
prefix_caches=[
|
||||
[{} for _ in range(config.dit_config.num_layers)]
|
||||
for _ in range(sample_count)
|
||||
],
|
||||
)
|
||||
sched = self.scheduler.config
|
||||
batch.extra["qwen21_mu"] = calculate_linear_shift(
|
||||
(batch.height // 16) * (batch.width // 16),
|
||||
base_seq_len=sched.get("base_image_seq_len", 256),
|
||||
max_seq_len=sched.get("max_image_seq_len", 4096),
|
||||
base_shift=sched.get("base_shift", 0.5),
|
||||
max_shift=sched.get("max_shift", 1.15),
|
||||
)
|
||||
return batch
|
||||
|
||||
|
||||
def prepare_qwen21_mu(batch, server_args):
|
||||
return "mu", batch.extra["qwen21_mu"]
|
||||
|
||||
|
||||
class QwenImage21DenoisingStage(DenoisingStage):
|
||||
def _predict_noise(
|
||||
self,
|
||||
current_model,
|
||||
latent_model_input,
|
||||
timestep,
|
||||
target_dtype,
|
||||
guidance,
|
||||
**kwargs,
|
||||
):
|
||||
caches = kwargs["prefix_caches"]
|
||||
if caches is not None and not caches[0][0]:
|
||||
# prefill is request-specific; graph replay must only see populated cache tensors
|
||||
return current_model(
|
||||
hidden_states=latent_model_input, timestep=timestep, **kwargs
|
||||
)
|
||||
return super()._predict_noise(
|
||||
current_model,
|
||||
latent_model_input,
|
||||
timestep,
|
||||
target_dtype,
|
||||
guidance,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -215,8 +215,10 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset(
|
||||
"minimaxai/minimax-h3",
|
||||
"qwen/qwen-image",
|
||||
"qwen/qwen-image-2512",
|
||||
"qwen/qwen-image-2.1",
|
||||
"qwen-image",
|
||||
"qwen-image-2512",
|
||||
"qwen-image-2.1",
|
||||
"tongyi-mai/z-image",
|
||||
"tongyi-mai/z-image-turbo",
|
||||
"zai-org/glm-image",
|
||||
@@ -236,6 +238,7 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset(
|
||||
"LongCatImagePipelineConfig",
|
||||
"MiniMaxH3PipelineConfig",
|
||||
"QwenImagePipelineConfig",
|
||||
"QwenImage21PipelineConfig",
|
||||
"SanaPipelineConfig",
|
||||
"SanaVideoPipelineConfig",
|
||||
"ZImagePipelineConfig",
|
||||
@@ -774,7 +777,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
logger.warning(
|
||||
"[Diffusion BCG] disabled for %s: only FLUX.1-dev, Ideogram-4, "
|
||||
"jdopensource/JoyAI-Echo, Lightricks/LTX-2, LongCat-Image, "
|
||||
"MiniMax-H3, Qwen/Qwen-Image, Qwen/Qwen-Image-2512, SANA1.5, "
|
||||
"MiniMax-H3, Qwen/Qwen-Image, Qwen/Qwen-Image-2512, "
|
||||
"Qwen/Qwen-Image-2.1, SANA1.5, "
|
||||
"SANA-Video, Tongyi-MAI/Z-Image/Z-Image-Turbo, and "
|
||||
"zai-org/GLM-Image are currently supported.",
|
||||
pipeline_config_name,
|
||||
|
||||
@@ -1131,6 +1131,26 @@ TWO_GPU_CASES = [
|
||||
ring_degree=2,
|
||||
),
|
||||
),
|
||||
# TODO: re-enable when the checkpoint is accessible to fork PR CI
|
||||
# DiffusionTestCase(
|
||||
# "qwen_image21_t2i_tp2",
|
||||
# DiffusionServerArgs(
|
||||
# model_path="Qwen/Qwen-Image-2.1",
|
||||
# tp_size=2,
|
||||
# ulysses_degree=1,
|
||||
# ring_degree=1,
|
||||
# ),
|
||||
# replace(
|
||||
# T2I_sampling_params,
|
||||
# output_size="1024x1024",
|
||||
# output_format="png",
|
||||
# extras={"num_inference_steps": 40, "guidance_scale": 1, "seed": 42},
|
||||
# ),
|
||||
# perf_repeat_requests=2,
|
||||
# run_perf_check=False,
|
||||
# run_component_accuracy_check=False,
|
||||
# run_t2v_input_reference_check=False,
|
||||
# ),
|
||||
DiffusionTestCase(
|
||||
"qwen_image_t2i_2_gpus_extra_high",
|
||||
DiffusionServerArgs(
|
||||
|
||||
@@ -179,6 +179,14 @@
|
||||
"runtime_peak_allocated_mb": 44855.0,
|
||||
"estimated_full_test_time_s": 65.6
|
||||
},
|
||||
"qwen_image21_t2i_tp2": {
|
||||
"stages_ms": {},
|
||||
"denoise_step_ms": {},
|
||||
"expected_e2e_ms": 0.0,
|
||||
"expected_avg_denoise_ms": 0.0,
|
||||
"expected_median_denoise_ms": 0.0,
|
||||
"estimated_full_test_time_s": 300.0
|
||||
},
|
||||
"qwen_image_t2i_2_gpus_extra_high": {
|
||||
"stages_ms": {},
|
||||
"denoise_step_ms": {},
|
||||
|
||||
@@ -1437,8 +1437,9 @@ Pinned revision used by this check: {SGL_TEST_FILES_CI_DATA_REVISION}
|
||||
assert model["object"] == "model", (
|
||||
f"Expected object='model', got {model.get('object')}"
|
||||
)
|
||||
assert model["id"] == case.server_args.model_path, (
|
||||
f"Model ID mismatch: expected {case.server_args.model_path}, got {model['id']}"
|
||||
expected_model_id = case.expected_model_id or case.server_args.model_path
|
||||
assert model["id"] == expected_model_id, (
|
||||
f"Model ID mismatch: expected {expected_model_id}, got {model['id']}"
|
||||
)
|
||||
|
||||
# Verify extended diffusion-specific fields
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Opt-in full-checkpoint tests until nightly runners can access the weights.
|
||||
|
||||
Set SGLANG_QWEN_IMAGE21_TEST_MODEL to an authorized model directory and
|
||||
SGLANG_QWEN_IMAGE21_TEST_IMAGE to a reference PNG to include editing.
|
||||
"""
|
||||
|
||||
import io
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from sglang.multimodal_gen.test.server.test_server_common import ( # noqa: F401
|
||||
DiffusionServerBase,
|
||||
diffusion_server,
|
||||
)
|
||||
from sglang.multimodal_gen.test.server.testcase_configs import (
|
||||
DiffusionSamplingParams,
|
||||
DiffusionServerArgs,
|
||||
DiffusionTestCase,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not os.environ.get("SGLANG_QWEN_IMAGE21_TEST_MODEL"),
|
||||
reason="requires an authorized Qwen-Image 2.1 checkpoint",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(params=["generation", "edit", "alpha"])
|
||||
def case(request):
|
||||
mode = request.param
|
||||
image = None
|
||||
prompt = "A red ceramic teapot on a wooden table beside a window."
|
||||
if mode == "edit":
|
||||
image_path = os.environ.get("SGLANG_QWEN_IMAGE21_TEST_IMAGE")
|
||||
if not image_path:
|
||||
pytest.skip("set SGLANG_QWEN_IMAGE21_TEST_IMAGE for the editing test")
|
||||
image = Path(image_path)
|
||||
assert image.is_file(), f"Reference image does not exist: {image}"
|
||||
prompt = "Change the teapot to blue, keeping its shape and the scene unchanged."
|
||||
elif mode == "alpha":
|
||||
prompt = (
|
||||
"A single fluffy orange cat sitting, full body, isolated on a transparent "
|
||||
"background. A clean cutout with an alpha channel, transparent outside "
|
||||
"the cat, no floor, no shadow, no background."
|
||||
)
|
||||
return DiffusionTestCase(
|
||||
f"qwen_image21_{mode}",
|
||||
DiffusionServerArgs(
|
||||
model_path=os.environ["SGLANG_QWEN_IMAGE21_TEST_MODEL"],
|
||||
modality="image",
|
||||
extras=[
|
||||
"--model-id Qwen-Image-2.1",
|
||||
"--performance-mode speed",
|
||||
"--attention-backend torch_sdpa",
|
||||
],
|
||||
),
|
||||
DiffusionSamplingParams(
|
||||
prompt=prompt,
|
||||
image_path=image,
|
||||
output_size="1024x1024",
|
||||
output_format="png",
|
||||
extras={"num_inference_steps": 40, "guidance_scale": 1, "seed": 42},
|
||||
),
|
||||
perf_repeat_requests=2,
|
||||
run_perf_check=False,
|
||||
run_consistency_check=False,
|
||||
run_component_accuracy_check=False,
|
||||
expected_model_id="Qwen-Image-2.1",
|
||||
run_t2v_input_reference_check=False,
|
||||
)
|
||||
|
||||
|
||||
class TestQwenImage21Server(DiffusionServerBase):
|
||||
def run_and_collect(self, ctx, case_id, generate_fn, collect_perf=True):
|
||||
record, content = super().run_and_collect(
|
||||
ctx, case_id, generate_fn, collect_perf
|
||||
)
|
||||
with Image.open(io.BytesIO(content)) as image:
|
||||
assert image.mode == "RGBA"
|
||||
assert image.size == (1024, 1024)
|
||||
if case_id.endswith("_alpha"):
|
||||
alpha = np.asarray(image.getchannel("A"))
|
||||
assert alpha.min() == 0 and alpha.max() == 255
|
||||
assert np.mean(alpha <= 5) > 0.4
|
||||
return record, content
|
||||
@@ -323,6 +323,7 @@ class DiffusionTestCase:
|
||||
run_consistency_check: bool = True
|
||||
run_component_accuracy_check: bool = True
|
||||
run_models_api_check: bool = True
|
||||
expected_model_id: str | None = None
|
||||
run_t2v_input_reference_check: bool = True
|
||||
run_lora_basic_api_check: bool = False
|
||||
run_lora_dynamic_load_check: bool = False
|
||||
|
||||
@@ -40,7 +40,7 @@ logger = init_logger(__name__)
|
||||
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
||||
# publish.
|
||||
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
||||
SGL_TEST_FILES_CI_DATA_REVISION = "0b9d7313c6bd31795fe6531a61ac45c89d9ed78e"
|
||||
SGL_TEST_FILES_CI_DATA_REVISION = "90a87cce5cdef73a9cd461f6d611ac66becef835"
|
||||
|
||||
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
||||
# when it's regenerated on its own cadence.
|
||||
@@ -171,6 +171,7 @@ DEFAULT_COSMOS3_NANO_MODEL_NAME_FOR_TEST = "nvidia/Cosmos3-Nano"
|
||||
|
||||
# Qwen image generation models
|
||||
DEFAULT_QWEN_IMAGE_MODEL_NAME_FOR_TEST = "Qwen/Qwen-Image"
|
||||
DEFAULT_QWEN_IMAGE_21_MODEL_NAME_FOR_TEST = "Qwen/Qwen-Image-2.1"
|
||||
DEFAULT_QWEN_IMAGE_2512_MODEL_NAME_FOR_TEST = "Qwen/Qwen-Image-2512"
|
||||
DEFAULT_QWEN_IMAGE_EDIT_MODEL_NAME_FOR_TEST = "Qwen/Qwen-Image-Edit"
|
||||
DEFAULT_QWEN_IMAGE_EDIT_2509_MODEL_NAME_FOR_TEST = "Qwen/Qwen-Image-Edit-2509"
|
||||
|
||||
@@ -357,6 +357,20 @@ class TestBuildCustomBlockAdapter(unittest.TestCase):
|
||||
self.assertEqual(adapter.forward_pattern, "Pattern_3")
|
||||
self.assertFalse(adapter.has_separate_cfg)
|
||||
|
||||
def test_native_qwen21_adapter_overrides_generic_family_match(self):
|
||||
module = _import_module_with_stub()
|
||||
module.BlockAdapterRegister.supported = True
|
||||
transformer = _make_transformer("QwenImage21Transformer2DModel")
|
||||
transformer.transformer_blocks = ["block_0"]
|
||||
config = module.CacheDitConfig(enabled=True, num_inference_steps=6)
|
||||
|
||||
module.enable_cache_on_transformer(transformer, config, has_separate_cfg=True)
|
||||
|
||||
adapter = module.cache_dit.enable_calls[0]["target"]
|
||||
self.assertEqual(adapter.forward_pattern, "Pattern_3")
|
||||
self.assertTrue(adapter.has_separate_cfg)
|
||||
self.assertIs(adapter.blocks, transformer.transformer_blocks)
|
||||
|
||||
def test_custom_adapter_is_retained_until_disable(self):
|
||||
module = _import_module_with_stub()
|
||||
module.BlockAdapterRegister.supported = False
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.disaggregation.roles import (
|
||||
RoleType,
|
||||
filter_modules_for_role,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
|
||||
SchedulerDisaggMixin,
|
||||
extract_transfer_fields,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.disaggregation.transport.codec import (
|
||||
pack_tensors,
|
||||
unpack_tensors,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image21 import build_layout
|
||||
from sglang.multimodal_gen.runtime.pipelines.qwen_image21 import QwenImage21Pipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
|
||||
|
||||
def test_qwen21_disagg_encoder_loads_condition_vae():
|
||||
pipeline = object.__new__(QwenImage21Pipeline)
|
||||
modules = filter_modules_for_role(
|
||||
pipeline._required_config_modules,
|
||||
RoleType.ENCODER,
|
||||
extra_allowed_modules=pipeline._get_extra_allowed_modules_for_role(
|
||||
RoleType.ENCODER, "ti2i"
|
||||
),
|
||||
)
|
||||
assert set(modules) == {"processor", "text_encoder", "vae", "scheduler"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("edit", [False, True])
|
||||
def test_qwen21_conditioning_survives_disagg_transfer(edit):
|
||||
slots = [False, True, False] if edit else [False, False, False]
|
||||
shapes = [(1, 2, 4), (1, 4, 4)] if edit else [(1, 4, 4)]
|
||||
condition = dict(
|
||||
layouts=[build_layout(slots, shapes, (8, 12, 12), "cpu")],
|
||||
condition_latents=torch.randn(1, 8, 4) if edit else None,
|
||||
prefix_caches=[[{}, {}]],
|
||||
)
|
||||
req = Req(request_id="qwen21-transfer", prompt="test")
|
||||
req.extra = dict(qwen21_positive=condition, qwen21_negative=condition, mu=0.7)
|
||||
req.extra["_local"] = object()
|
||||
tensors, scalars = extract_transfer_fields(req)
|
||||
metadata, buffers = pack_tensors(tensors, scalars)
|
||||
received, scalars = unpack_tensors([metadata, *[w._view for w in buffers]])
|
||||
rebuilt = SchedulerDisaggMixin._build_disagg_req(None, scalars, received)
|
||||
assert "_local" not in rebuilt.extra
|
||||
assert rebuilt.extra["mu"] == 0.7
|
||||
for name in ("qwen21_positive", "qwen21_negative"):
|
||||
restored = rebuilt.extra[name]
|
||||
for key, expected in condition["layouts"][0].items():
|
||||
actual = restored["layouts"][0][key]
|
||||
if isinstance(expected, torch.Tensor):
|
||||
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
|
||||
else:
|
||||
assert actual == expected
|
||||
assert restored["prefix_caches"] == [[{}, {}]]
|
||||
if edit:
|
||||
torch.testing.assert_close(
|
||||
restored["condition_latents"],
|
||||
condition["condition_latents"],
|
||||
atol=0,
|
||||
rtol=0,
|
||||
)
|
||||
else:
|
||||
assert restored["condition_latents"] is None
|
||||
@@ -1,5 +1,6 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
@@ -26,6 +27,28 @@ from sglang.srt.models.qwen3_vl import (
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
@pytest.mark.parametrize("dim", [36, 40, 64])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])
|
||||
@pytest.mark.parametrize("recompute_on_device_change", [False, True])
|
||||
def test_vision_rope_device_transfer(dim, dtype, recompute_on_device_change):
|
||||
with torch.device("cpu"):
|
||||
transferred = Qwen3VLVisionRotaryEmbedding(dim).to(dtype=dtype)
|
||||
assert transferred.recompute_on_device_change is False
|
||||
transferred.recompute_on_device_change = recompute_on_device_change
|
||||
expected_cpu = transferred(64).clone()
|
||||
with torch.device("cuda"):
|
||||
resident = Qwen3VLVisionRotaryEmbedding(dim).to(dtype=dtype)
|
||||
expected_cuda = resident(64)
|
||||
|
||||
transferred.cuda()
|
||||
expected = expected_cuda if recompute_on_device_change else expected_cpu.cuda()
|
||||
torch.testing.assert_close(transferred(64), expected, atol=0, rtol=0)
|
||||
torch.testing.assert_close(transferred(64), expected, atol=0, rtol=0)
|
||||
transferred.cpu()
|
||||
torch.testing.assert_close(transferred(64), expected_cpu, atol=0, rtol=0)
|
||||
|
||||
|
||||
def test_native_vision_layout_matches_qwen3_merge_order():
|
||||
grid_thw = torch.tensor([[1, 4, 6], [2, 2, 4]])
|
||||
|
||||
@@ -191,7 +214,7 @@ def test_qwen3vl_ties_lm_head_to_input_embeddings():
|
||||
_fsdp_shard_conditions=[],
|
||||
stacked_params_mapping=[],
|
||||
)
|
||||
config = SimpleNamespace(arch_config=arch_config)
|
||||
config = SimpleNamespace(arch_config=arch_config, quant_config=None)
|
||||
|
||||
with get_parallel().override(tp_size=1, tp_rank=0):
|
||||
model = Qwen3VLForConditionalGeneration(config)
|
||||
@@ -214,3 +237,25 @@ def test_qwen3_multimodal_encoders_layerwise_offload_vision_blocks():
|
||||
condition.__name__ == "is_block"
|
||||
for condition in Qwen3VLArchConfig()._fsdp_shard_conditions
|
||||
)
|
||||
|
||||
|
||||
def test_vision_position_interpolation_preserves_bf16_rounding():
|
||||
model = Qwen3VLVisionTransformer.__new__(Qwen3VLVisionTransformer)
|
||||
nn.Module.__init__(model)
|
||||
model.num_grid_per_side = 2
|
||||
model.spatial_merge_size = 2
|
||||
model.pos_embed = nn.Embedding.from_pretrained(
|
||||
torch.tensor(
|
||||
[[7.21875], [-3.359375], [2.078125], [-1.1171875]], dtype=torch.bfloat16
|
||||
)
|
||||
)
|
||||
model.fp32_position_interpolation = False
|
||||
positions = model._interpolate_position_embeddings(torch.tensor([[1, 4, 4]]))
|
||||
# (1, 1) has corner weights 4/9, 2/9, 2/9, 1/9 in merge order
|
||||
assert positions.dtype == torch.bfloat16
|
||||
assert positions[3, 0].item() == 2.8125
|
||||
model.fp32_position_interpolation = True
|
||||
assert (
|
||||
model._interpolate_position_embeddings(torch.tensor([[1, 4, 4]])).dtype
|
||||
== torch.float32
|
||||
)
|
||||
|
||||
@@ -0,0 +1,295 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from contextlib import nullcontext
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from PIL import Image
|
||||
from transformers import BatchFeature
|
||||
|
||||
from sglang.multimodal_gen.configs.models.dits.qwenimage21 import (
|
||||
QwenImage21ArchConfig,
|
||||
QwenImage21DitConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.vaes.qwenimage21 import (
|
||||
QwenImage21VAEArchConfig,
|
||||
QwenImage21VAEConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import (
|
||||
QwenImage21PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.registry import _get_config_info
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import (
|
||||
ResidencyState,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency_strategies import (
|
||||
ComponentOffloadStrategy,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image21 import build_layout
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import (
|
||||
Qwen3VLVisionRotaryEmbedding,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage21 import (
|
||||
AutoencoderKLQwenImage21,
|
||||
QwenImage21RMS_norm,
|
||||
_patchify,
|
||||
_unpatchify,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
|
||||
InputValidationStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image21 import (
|
||||
QwenImage21EncodingStage,
|
||||
QwenImage21InputValidationStage,
|
||||
collapse_image_slots,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", ["edit", ""])
|
||||
@pytest.mark.parametrize("image_count", [0, 1, 2])
|
||||
def test_prompt_conditioning_uses_training_template_and_pre_norm(prompt, image_count):
|
||||
hidden = torch.arange(24).reshape(1, 6, 4).float()
|
||||
inputs = BatchFeature(
|
||||
data={
|
||||
"input_ids": torch.tensor([[1, 2, 99, 99, 3, 0]]),
|
||||
"attention_mask": torch.tensor([[1, 1, 1, 1, 1, 0]]),
|
||||
}
|
||||
)
|
||||
processor = Mock(return_value=inputs)
|
||||
processor.tokenizer.convert_tokens_to_ids.return_value = 99
|
||||
processor.apply_chat_template.return_value = [[1]]
|
||||
encoder = Mock(return_value=SimpleNamespace(hidden_states=(hidden,)))
|
||||
stage = QwenImage21EncodingStage(encoder, processor, None, None)
|
||||
stage.use_declared_component = Mock(return_value=nullcontext(encoder))
|
||||
images = [Image.new("RGBA", (2, 1), (12, 34, 56, 0)) for _ in range(image_count)]
|
||||
for image in images:
|
||||
image.putpixel((1, 0), (12, 34, 56, 255))
|
||||
actual, slots = stage.encode_prompt(prompt, images, "cpu")
|
||||
torch.testing.assert_close(actual, hidden[0, [1, 2, 4]])
|
||||
assert slots.tolist() == [False, True, False]
|
||||
encoder.model.language_model.norm.assert_not_called()
|
||||
assert encoder.model.visual.fp32_position_interpolation is False
|
||||
assert encoder.model.visual.rotary_pos_emb.recompute_on_device_change is True
|
||||
kwargs = processor.call_args.kwargs
|
||||
prefix = " ".join(
|
||||
f"<image{i + 1}><|vision_start|><|image_pad|><|vision_end|>"
|
||||
for i in range(image_count)
|
||||
)
|
||||
assert kwargs["text"] == [
|
||||
"<|im_start|>system\nComprehend and analyze the provided prompt.<|im_end|>\n"
|
||||
f"<|im_start|>user\n{prefix}{prompt or ' '}<|im_end|>\n<|im_start|>assistant\n"
|
||||
]
|
||||
assert kwargs["padding_side"] == "left"
|
||||
for image in kwargs.get("images", []):
|
||||
assert image.mode == "RGB"
|
||||
assert image.getpixel((0, 0)) == (255, 255, 255)
|
||||
assert image.getpixel((1, 0)) == (12, 34, 56)
|
||||
for image in images:
|
||||
assert image.mode == "RGBA"
|
||||
assert image.getpixel((0, 0)) == (12, 34, 56, 0)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
def test_encoder_component_offload_preserves_loaded_dtypes(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"sglang.multimodal_gen.runtime.managers.memory_managers."
|
||||
"component_residency_strategies.get_local_torch_device",
|
||||
lambda: torch.device("cuda", torch.cuda.current_device()),
|
||||
)
|
||||
encoder = torch.nn.Module()
|
||||
encoder.model = torch.nn.Module()
|
||||
encoder.model.visual = torch.nn.Module()
|
||||
encoder.model.visual.rotary_pos_emb = Qwen3VLVisionRotaryEmbedding(36)
|
||||
with torch.device("cuda"):
|
||||
expected_rope = Qwen3VLVisionRotaryEmbedding(36)(64)
|
||||
encoder.register_parameter(
|
||||
"embedding", torch.nn.Parameter(torch.ones(2, dtype=torch.bfloat16))
|
||||
)
|
||||
encoder.register_parameter(
|
||||
"weight",
|
||||
torch.nn.Parameter(
|
||||
torch.tensor([0.25, -0.5]).to(torch.float8_e4m3fn), requires_grad=False
|
||||
),
|
||||
)
|
||||
frequencies = torch.tensor([1.0 / 3, 1.0 / 7])
|
||||
encoder.register_buffer("inv_freq", frequencies.clone())
|
||||
processor = Mock()
|
||||
processor.apply_chat_template.return_value = [[1]]
|
||||
stage = QwenImage21EncodingStage(encoder, processor, None, None)
|
||||
use = stage.component_uses(None, "conditioning")[0]
|
||||
strategy = ComponentOffloadStrategy()
|
||||
state = ResidencyState(batch_is_warmup=False)
|
||||
weight_bytes = encoder.weight.view(torch.uint8).clone()
|
||||
|
||||
for _ in range(2):
|
||||
strategy.prefetch_for_use(encoder, use, state)
|
||||
strategy.wait_for_use(encoder, use, state)
|
||||
assert encoder.embedding.device.type == "cuda"
|
||||
assert encoder.embedding.dtype == torch.bfloat16
|
||||
assert encoder.weight.dtype == torch.float8_e4m3fn
|
||||
assert encoder.inv_freq.dtype == torch.float32
|
||||
torch.testing.assert_close(encoder.inv_freq.cpu(), frequencies, atol=0, rtol=0)
|
||||
torch.testing.assert_close(
|
||||
encoder.model.visual.rotary_pos_emb(64), expected_rope, atol=0, rtol=0
|
||||
)
|
||||
assert torch.equal(encoder.weight.view(torch.uint8).cpu(), weight_bytes)
|
||||
strategy.finish_use(encoder, use, state)
|
||||
torch.cuda.synchronize()
|
||||
assert encoder.embedding.device.type == "cpu"
|
||||
|
||||
|
||||
def test_condition_slots_expand_to_actual_latent_grid():
|
||||
hidden = torch.randn(22, 8)
|
||||
ids = torch.tensor([1, 2] + [99] * 16 + [3, 4, 5, 6])
|
||||
collapsed, slots = collapse_image_slots(hidden, ids, 99)
|
||||
assert collapsed.shape == (7, 8)
|
||||
layout = build_layout(slots.tolist(), [(1, 4, 8), (1, 2, 2)], (4, 6, 6), "cpu")
|
||||
assert len(layout["image_indices"]) == 32
|
||||
assert len(layout["prefix_rope"]) == 38
|
||||
assert layout["segments"] == ((0, 2, False), (2, 34, True), (34, 38, False))
|
||||
torch.testing.assert_close(collapsed[slots][0], hidden[2])
|
||||
|
||||
|
||||
def test_adjacent_image_slots_stay_distinct():
|
||||
layout = build_layout(
|
||||
[False, True, True, False], [(1, 2, 2), (1, 4, 2), (1, 2, 2)], (4, 6, 6), "cpu"
|
||||
)
|
||||
assert layout["segments"] == (
|
||||
(0, 1, False),
|
||||
(1, 5, True),
|
||||
(5, 13, True),
|
||||
(13, 14, False),
|
||||
)
|
||||
with pytest.raises(ValueError, match="slots"):
|
||||
build_layout([False], [(1, 2, 2), (1, 2, 2)], (4, 6, 6), "cpu")
|
||||
|
||||
|
||||
def test_latent_pack_decode_contract():
|
||||
config = QwenImage21PipelineConfig()
|
||||
batch = SimpleNamespace(
|
||||
height=64, width=96, extra={"qwen21_positive": {}, "qwen21_negative": {}}
|
||||
)
|
||||
shape = config.prepare_latent_shape(batch, 2, 1)
|
||||
x = torch.arange(torch.tensor(shape).prod()).reshape(shape)
|
||||
packed = config.maybe_pack_latents(x, 2, batch)
|
||||
assert packed.shape == (2, 24, 64)
|
||||
decoded = config.post_denoising_loop(packed, batch)
|
||||
assert not batch.extra
|
||||
torch.testing.assert_close(decoded[:, :, 0], x[:, 0])
|
||||
scale, shift = config.get_decode_scale_and_shift("cpu", torch.float32, None)
|
||||
torch.testing.assert_close(
|
||||
(decoded.float() - shift) * scale / scale + shift, decoded.float()
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("channels", [3, 4])
|
||||
@pytest.mark.parametrize("tiling", [False, True])
|
||||
def test_native_vae_roundtrip_shapes_and_checkpoint_names(channels, tiling):
|
||||
ac = QwenImage21VAEArchConfig(
|
||||
base_dim=4,
|
||||
decoder_base_dim=4,
|
||||
z_dim=4,
|
||||
dim_mult=(1, 2, 4, 4, 4),
|
||||
num_res_blocks=1,
|
||||
temperal_downsample=(False, False, False, False),
|
||||
in_channels=channels,
|
||||
out_channels=channels,
|
||||
)
|
||||
model = AutoencoderKLQwenImage21(QwenImage21VAEConfig(arch_config=ac)).eval()
|
||||
assert not model.use_tiling
|
||||
assert ac.scale_factor_spatial == ac.spatial_compression_ratio == 16
|
||||
model.use_tiling = tiling
|
||||
model.use_parallel_tiling = False
|
||||
model.tile_sample_min_height = model.tile_sample_min_width = 32
|
||||
model.tile_sample_stride_height = model.tile_sample_stride_width = 16
|
||||
with torch.no_grad():
|
||||
latent = model.encode(torch.randn(1, channels, 1, 32, 64)).mode()
|
||||
assert latent.shape == (1, 4, 1, 2, 4)
|
||||
output = model.decode(latent)
|
||||
assert output.shape == (1, channels, 1, 32, 64)
|
||||
assert model.state_dict()["encoder.conv_in.weight"].ndim == 4
|
||||
x = torch.randn(2, 3, 1, 8, 12)
|
||||
torch.testing.assert_close(_unpatchify(_patchify(x, 2), 2), x)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
|
||||
def test_vae_rms_norm_normalizes_in_float32(dtype):
|
||||
norm = QwenImage21RMS_norm(8, images=False).to(dtype)
|
||||
x = torch.linspace(-60000, 60000, 256).reshape(1, 8, 1, 4, 8).to(dtype)
|
||||
expected = (
|
||||
torch.nn.functional.normalize(x.float(), dim=1).to(dtype)
|
||||
* norm.scale
|
||||
* norm.gamma
|
||||
)
|
||||
torch.testing.assert_close(norm(x), expected, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tiling", [False, True])
|
||||
def test_condition_pixels_match_reference_preprocessing(monkeypatch, tiling):
|
||||
module = "sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.qwen_image21"
|
||||
monkeypatch.setattr(f"{module}.get_local_torch_device", lambda: torch.device("cpu"))
|
||||
monkeypatch.setattr(f"{module}.set_forward_context", lambda **kwargs: nullcontext())
|
||||
image = Image.frombytes("RGBA", (32, 32), bytes(range(256)) * 16)
|
||||
vae = Mock()
|
||||
vae.encode.return_value.mode.return_value = torch.zeros(1, 64, 1, 2, 2)
|
||||
processor = Mock()
|
||||
processor.apply_chat_template.return_value = [[1]]
|
||||
stage = QwenImage21EncodingStage(Mock(), processor, vae, SimpleNamespace(config={}))
|
||||
stage.use_declared_component = Mock(return_value=nullcontext(vae))
|
||||
stage.encode_prompt = Mock(
|
||||
return_value=(torch.zeros(3, 8), torch.tensor([False, True, False]))
|
||||
)
|
||||
batch = SimpleNamespace(
|
||||
height=32,
|
||||
width=32,
|
||||
condition_image=image,
|
||||
prompt="edit",
|
||||
negative_prompt=None,
|
||||
num_outputs_per_prompt=1,
|
||||
do_classifier_free_guidance=False,
|
||||
extra={},
|
||||
)
|
||||
stage.forward(
|
||||
batch,
|
||||
SimpleNamespace(pipeline_config=QwenImage21PipelineConfig(vae_tiling=tiling)),
|
||||
)
|
||||
assert vae.use_tiling is tiling
|
||||
expected = VaeImageProcessor(vae_scale_factor=16).preprocess(image).unsqueeze(2)
|
||||
actual = vae.encode.call_args.args[0]
|
||||
torch.testing.assert_close(actual, expected.bfloat16(), atol=0, rtol=0)
|
||||
assert actual.stride() == expected.stride()
|
||||
|
||||
|
||||
def test_condition_image_loading_preserves_alpha(tmp_path):
|
||||
path = tmp_path / "condition.png"
|
||||
Image.new("RGBA", (32, 32), (12, 34, 56, 78)).save(path)
|
||||
image = QwenImage21InputValidationStage().load_condition_image(str(path))
|
||||
assert image.mode == "RGBA"
|
||||
assert image.getpixel((0, 0)) == (12, 34, 56, 78)
|
||||
assert InputValidationStage().load_condition_image(str(path)).mode == "RGB"
|
||||
|
||||
|
||||
def test_architecture_derived_dimensions():
|
||||
config = QwenImage21DitConfig(
|
||||
arch_config=QwenImage21ArchConfig(num_attention_heads=2, attention_head_dim=16)
|
||||
)
|
||||
assert config.hidden_size == 32
|
||||
|
||||
|
||||
def test_registry_routes_local_checkpoint_and_preserves_legacy():
|
||||
assert (
|
||||
_get_config_info("Qwen/Qwen-Image-2.1").pipeline_config_cls
|
||||
is QwenImage21PipelineConfig
|
||||
)
|
||||
assert (
|
||||
_get_config_info(
|
||||
"/models/private", model_id="Qwen-Image-2.1"
|
||||
).pipeline_config_cls
|
||||
is QwenImage21PipelineConfig
|
||||
)
|
||||
assert (
|
||||
_get_config_info("Qwen/Qwen-Image").pipeline_config_cls
|
||||
is not QwenImage21PipelineConfig
|
||||
)
|
||||
@@ -0,0 +1,270 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Request-scoped prefix KV and graph replay regression tests; no checkpoint needed."""
|
||||
|
||||
from copy import deepcopy
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from diffusers.models.normalization import RMSNorm as ReferenceRMSNorm
|
||||
from safetensors.torch import save_file
|
||||
|
||||
from sglang.kernels.ops.diffusion import BitExactFusionGate
|
||||
from sglang.multimodal_gen.configs.models.dits.qwenimage21 import (
|
||||
QwenImage21ArchConfig,
|
||||
QwenImage21DitConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.pipeline_configs.qwen_image21 import (
|
||||
QwenImage21PipelineConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
|
||||
DiffusionBreakableCudaGraphRunner,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||
from sglang.multimodal_gen.runtime.models.dits import qwen_image21 as model_module
|
||||
from sglang.multimodal_gen.runtime.models.dits.qwen_image21 import (
|
||||
QwenImage21Transformer2DModel,
|
||||
build_layout,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines.qwen_image21 import QwenImage21Pipeline
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import (
|
||||
ComposedPipelineBase,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import (
|
||||
ServerArgs,
|
||||
set_global_server_args,
|
||||
)
|
||||
from sglang.multimodal_gen.test.single_test_file.component_accuracy.utils import (
|
||||
ensure_distributed_env_defaults,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model():
|
||||
config = QwenImage21DitConfig(
|
||||
arch_config=QwenImage21ArchConfig(
|
||||
in_channels=4,
|
||||
out_channels=4,
|
||||
num_layers=3,
|
||||
num_attention_heads=4,
|
||||
attention_head_dim=32,
|
||||
context_in_dim=16,
|
||||
mlp_ratio=2,
|
||||
axes_dims_rope=(8, 12, 12),
|
||||
)
|
||||
)
|
||||
args = ServerArgs(
|
||||
model_path="Qwen/Qwen-Image-2.1",
|
||||
num_gpus=1,
|
||||
pipeline_config=QwenImage21PipelineConfig(dit_config=config),
|
||||
attention_backend="torch_sdpa",
|
||||
)
|
||||
set_global_server_args(args)
|
||||
if not model_parallel_is_initialized():
|
||||
ensure_distributed_env_defaults()
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
torch.manual_seed(42)
|
||||
model = QwenImage21Transformer2DModel(config, {}).cuda().eval()
|
||||
# Parallel linear layers allocate empty weights for checkpoint loading.
|
||||
for name, param in model.named_parameters():
|
||||
torch.nn.init.normal_(param, std=0.02)
|
||||
if name.endswith(("norm_q.weight", "norm_k.weight")):
|
||||
torch.nn.init.ones_(param)
|
||||
return model
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bf16_model(model):
|
||||
# parallel modules own process groups and cannot be deep-copied
|
||||
config = QwenImage21DitConfig(arch_config=model.config)
|
||||
result = QwenImage21Transformer2DModel(config, {}).cuda().bfloat16().eval()
|
||||
result.load_state_dict(model.state_dict())
|
||||
return result
|
||||
|
||||
|
||||
def inputs(seed, edit):
|
||||
torch.manual_seed(seed)
|
||||
slots = [False] * 3 + ([True, False, False] if edit else [])
|
||||
shapes = ([(1, 2, 4)] if edit else []) + [(1, 4, 4)]
|
||||
return dict(
|
||||
hidden_states=torch.randn(1, 16, 4, device="cuda"),
|
||||
encoder_hidden_states=torch.randn(1, len(slots), 16, device="cuda"),
|
||||
condition_latents=torch.randn(1, 8, 4, device="cuda") if edit else None,
|
||||
layouts=[build_layout(slots, shapes, (8, 12, 12), "cuda")],
|
||||
prefix_caches=[[{} for _ in range(3)]],
|
||||
timestep=torch.tensor([700.0], device="cuda"),
|
||||
)
|
||||
|
||||
|
||||
def test_bf16_qk_norm_matches_reference(model):
|
||||
norm = deepcopy(model.transformer_blocks[0].attn.norm_q).bfloat16()
|
||||
reference = ReferenceRMSNorm(32, eps=1e-6).cuda().bfloat16()
|
||||
weight = torch.linspace(0.3, 1.7, 32, device="cuda", dtype=torch.bfloat16)
|
||||
x = torch.randn(2, 8, 4, 32, device="cuda", dtype=torch.bfloat16)
|
||||
with torch.no_grad():
|
||||
norm.weight.copy_(weight)
|
||||
reference.weight.copy_(weight)
|
||||
torch.testing.assert_close(norm(x), reference(x), atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("merge_mode", ["dynamic", "merge"])
|
||||
@torch.no_grad()
|
||||
def test_diffusers_lora_matches_weight_delta_and_restores_base(
|
||||
model, tmp_path, monkeypatch, merge_mode
|
||||
):
|
||||
# Reuse loaded native components, then exercise the real adapter loader.
|
||||
monkeypatch.setattr(ComposedPipelineBase, "__init__", lambda self: None)
|
||||
pipeline = object.__new__(QwenImage21Pipeline)
|
||||
config = QwenImage21DitConfig(arch_config=model.config)
|
||||
pipeline.server_args = ServerArgs(
|
||||
model_path="Qwen/Qwen-Image-2.1",
|
||||
num_gpus=1,
|
||||
pipeline_config=QwenImage21PipelineConfig(dit_config=config),
|
||||
attention_backend="torch_sdpa",
|
||||
)
|
||||
set_global_server_args(pipeline.server_args)
|
||||
actual_model = QwenImage21Transformer2DModel(config, {}).cuda().eval()
|
||||
reference = QwenImage21Transformer2DModel(config, {}).cuda().eval()
|
||||
for loaded in (actual_model, reference):
|
||||
loaded.load_state_dict(model.state_dict())
|
||||
pipeline.modules = {"transformer": actual_model}
|
||||
pipeline.__init__()
|
||||
weights = {}
|
||||
for name in ("transformer_blocks.0.attn.to_q", "transformer_blocks.0.img_mlp.out"):
|
||||
layer = reference.get_submodule(name)
|
||||
a = torch.randn(2, layer.weight.shape[1], device="cuda") * 0.2
|
||||
b = torch.randn(layer.weight.shape[0], 2, device="cuda") * 0.2
|
||||
weights[f"transformer.{name}.lora_A.weight"] = a.cpu()
|
||||
weights[f"transformer.{name}.lora_B.weight"] = b.cpu()
|
||||
layer.weight.add_(b @ a)
|
||||
adapter = tmp_path / "adapter.safetensors"
|
||||
save_file(weights, str(adapter))
|
||||
kwargs = dict(inputs(5, False), prefix_caches=None)
|
||||
with set_forward_context(None, None):
|
||||
baseline = actual_model(**kwargs)
|
||||
expected = reference(**kwargs)
|
||||
pipeline.set_lora(
|
||||
"test", str(adapter), target="transformer", merge_mode=merge_mode
|
||||
)
|
||||
assert pipeline.is_lora_effective("transformer")
|
||||
actual = actual_model(**kwargs)
|
||||
assert not torch.equal(actual, baseline)
|
||||
torch.testing.assert_close(actual, expected, atol=1e-5, rtol=1e-5)
|
||||
pipeline.unmerge_lora_weights("transformer")
|
||||
torch.testing.assert_close(actual_model(**kwargs), baseline, atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("edit", [False, True])
|
||||
def test_cached_prefix_matches_full_recomputation(model, edit):
|
||||
kwargs = inputs(5, edit)
|
||||
with torch.no_grad(), set_forward_context(None, None):
|
||||
model(**kwargs)
|
||||
prefix_length = kwargs["layouts"][0]["prefix_rope"].shape[0]
|
||||
for cache in kwargs["prefix_caches"][0]:
|
||||
for tensor in cache.values():
|
||||
assert tensor.shape[1] == prefix_length
|
||||
assert (
|
||||
tensor.untyped_storage().nbytes()
|
||||
== tensor.numel() * tensor.element_size()
|
||||
)
|
||||
keys = [layer["key"].clone() for layer in kwargs["prefix_caches"][0]]
|
||||
kwargs["timestep"].fill_(300.0)
|
||||
actual = model(**kwargs)
|
||||
expected = model(**dict(kwargs, prefix_caches=None))
|
||||
torch.testing.assert_close(actual, expected, atol=0, rtol=0)
|
||||
for key, cache in zip(keys, kwargs["prefix_caches"][0], strict=True):
|
||||
torch.testing.assert_close(key, cache["key"], atol=0, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("edit", [False, True])
|
||||
def test_graph_replay_uses_new_request_prefix(model, edit):
|
||||
first, second = inputs(5, edit), inputs(9, edit)
|
||||
runner = DiffusionBreakableCudaGraphRunner(model, torch.device("cuda"))
|
||||
try:
|
||||
with torch.no_grad(), set_forward_context(None, None):
|
||||
model(**first)
|
||||
assert runner.capture(**first)
|
||||
model(**second)
|
||||
expected = model(**second)
|
||||
actual = runner(**second)
|
||||
assert len(runner.entries) == 1
|
||||
torch.testing.assert_close(actual, expected, atol=1e-6, rtol=1e-6)
|
||||
finally:
|
||||
runner.reset()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("edit", [False, True])
|
||||
@torch.no_grad()
|
||||
def test_bf16_fusions_match_eager_prefill_and_cached_steps(
|
||||
bf16_model, edit, monkeypatch
|
||||
):
|
||||
actual_model = bf16_model
|
||||
kwargs = inputs(5, edit)
|
||||
for key in (
|
||||
"hidden_states",
|
||||
"encoder_hidden_states",
|
||||
"condition_latents",
|
||||
"timestep",
|
||||
):
|
||||
if kwargs[key] is not None:
|
||||
kwargs[key] = kwargs[key].bfloat16()
|
||||
reference_kwargs = deepcopy(kwargs)
|
||||
expected = []
|
||||
disabled = BitExactFusionGate("reference")
|
||||
disabled.disable()
|
||||
with monkeypatch.context() as reference, set_forward_context(None, None):
|
||||
reference.setattr(model_module, "_SILU_MUL_FUSION", disabled)
|
||||
reference.setattr(
|
||||
model_module,
|
||||
"residual_gate_add",
|
||||
lambda residual, update, gate: residual + gate * update,
|
||||
)
|
||||
for timestep in (700, 300, 10):
|
||||
reference_kwargs["timestep"].fill_(timestep)
|
||||
expected.append(actual_model(**reference_kwargs))
|
||||
|
||||
gate = BitExactFusionGate("test SiLU-mul")
|
||||
monkeypatch.setattr(model_module, "_SILU_MUL_FUSION", gate)
|
||||
with set_forward_context(None, None):
|
||||
for timestep, output in zip((700, 300, 10), expected, strict=True):
|
||||
kwargs["timestep"].fill_(timestep)
|
||||
torch.testing.assert_close(actual_model(**kwargs), output, atol=0, rtol=0)
|
||||
assert gate.verified and not gate.disabled
|
||||
for actual, reference in zip(
|
||||
kwargs["prefix_caches"][0], reference_kwargs["prefix_caches"][0], strict=True
|
||||
):
|
||||
for key in ("key", "value"):
|
||||
torch.testing.assert_close(actual[key], reference[key], atol=0, rtol=0)
|
||||
|
||||
runner = DiffusionBreakableCudaGraphRunner(actual_model, torch.device("cuda"))
|
||||
try:
|
||||
with set_forward_context(None, None):
|
||||
assert runner.capture(**kwargs)
|
||||
kwargs["hidden_states"].add_(0.1)
|
||||
expected = actual_model(**kwargs)
|
||||
torch.testing.assert_close(runner(**kwargs), expected, atol=0, rtol=0)
|
||||
finally:
|
||||
runner.reset()
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_silu_fusion_mismatch_restores_eager(bf16_model, monkeypatch):
|
||||
mlp = bf16_model.transformer_blocks[0].img_mlp
|
||||
x = torch.randn(1, 16, 128, device="cuda", dtype=torch.bfloat16)
|
||||
gate = BitExactFusionGate("test mismatch")
|
||||
monkeypatch.setattr(model_module, "_SILU_MUL_FUSION", gate)
|
||||
monkeypatch.setattr(
|
||||
model_module, "fused_silu_mul_bitexact", lambda a, b: torch.zeros_like(a)
|
||||
)
|
||||
with set_forward_context(None, None):
|
||||
expected = mlp.out(
|
||||
torch.nn.functional.silu(mlp.gate_layer(x)[0]) * mlp.proj(x)[0]
|
||||
)[0]
|
||||
torch.testing.assert_close(mlp(x), expected, atol=0, rtol=0)
|
||||
assert gate.disabled and not gate.verified
|
||||
torch.testing.assert_close(mlp(x), expected, atol=0, rtol=0)
|
||||
@@ -0,0 +1,171 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run with torchrun --standalone --nproc-per-node=2 -m pytest -q <this file>."""
|
||||
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from transformers.models.qwen3_vl.configuration_qwen3_vl import (
|
||||
Qwen3VLConfig as HFQwen3VLConfig,
|
||||
)
|
||||
|
||||
from sglang.multimodal_gen.configs.models.vaes.qwenimage21 import (
|
||||
QwenImage21VAEArchConfig,
|
||||
QwenImage21VAEConfig,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||
get_sp_group,
|
||||
get_tp_group,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
use_tensor_parallel_group,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||
from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import (
|
||||
Qwen3VLForConditionalGeneration,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.vaes.autoencoder_kl_qwenimage21 import (
|
||||
AutoencoderKLQwenImage21,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core import Req
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, set_global_server_args
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or int(os.environ.get("WORLD_SIZE", "1")) != 2,
|
||||
reason="requires two CUDA ranks launched by torchrun",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module", autouse=True)
|
||||
def distributed():
|
||||
args = ServerArgs(
|
||||
model_path="Qwen/Qwen-Image-2.1",
|
||||
num_gpus=2,
|
||||
tp_size=2,
|
||||
sp_degree=1,
|
||||
attention_backend="torch_sdpa",
|
||||
)
|
||||
set_global_server_args(args)
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=2, sp_size=1)
|
||||
matmul_tf32 = torch.backends.cuda.matmul.allow_tf32
|
||||
cudnn_tf32 = torch.backends.cudnn.allow_tf32
|
||||
torch.backends.cuda.matmul.allow_tf32 = False
|
||||
torch.backends.cudnn.allow_tf32 = False
|
||||
yield
|
||||
torch.backends.cuda.matmul.allow_tf32 = matmul_tf32
|
||||
torch.backends.cudnn.allow_tf32 = cudnn_tf32
|
||||
|
||||
|
||||
@pytest.mark.parametrize("edit", [False, True])
|
||||
@torch.no_grad()
|
||||
def test_encoder_tp_shards_weights_and_preserves_conditioning(edit):
|
||||
arch = HFQwen3VLConfig(
|
||||
text_config=dict(
|
||||
hidden_size=64,
|
||||
intermediate_size=128,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=4,
|
||||
num_key_value_heads=2,
|
||||
head_dim=16,
|
||||
vocab_size=32,
|
||||
pad_token_id=0,
|
||||
rope_scaling=dict(rope_type="default", mrope_section=[2, 3, 3]),
|
||||
),
|
||||
vision_config=dict(
|
||||
hidden_size=32,
|
||||
intermediate_size=64,
|
||||
depth=2,
|
||||
num_heads=4,
|
||||
patch_size=2,
|
||||
temporal_patch_size=1,
|
||||
in_channels=3,
|
||||
num_position_embeddings=16,
|
||||
spatial_merge_size=2,
|
||||
out_hidden_size=64,
|
||||
deepstack_visual_indexes=[],
|
||||
),
|
||||
image_token_id=8,
|
||||
video_token_id=9,
|
||||
vision_start_token_id=7,
|
||||
vision_end_token_id=6,
|
||||
)
|
||||
arch._fsdp_shard_conditions = []
|
||||
arch.stacked_params_mapping = []
|
||||
config = SimpleNamespace(arch_config=arch, quant_config=None)
|
||||
torch.manual_seed(42)
|
||||
with use_tensor_parallel_group(get_sp_group()):
|
||||
reference = Qwen3VLForConditionalGeneration(config).cuda().eval()
|
||||
for param in reference.parameters():
|
||||
torch.nn.init.normal_(param, std=0.02)
|
||||
reference.bind_encoder_tp_group(get_sp_group())
|
||||
with use_tensor_parallel_group(get_tp_group()):
|
||||
model = Qwen3VLForConditionalGeneration(config).cuda().eval()
|
||||
model.bind_encoder_tp_group(get_tp_group())
|
||||
model.load_weights(reference.state_dict().items())
|
||||
for layer in model.model.language_model.layers:
|
||||
assert layer.self_attn.q_proj.weight.shape == (32, 64)
|
||||
assert layer.mlp.gate_proj.weight.shape == (64, 64)
|
||||
tokens = [1, 7, 8, 8, 8, 8, 6, 3] if edit else [1, 2, 3, 4]
|
||||
inputs = dict(
|
||||
input_ids=torch.tensor([tokens], device="cuda"),
|
||||
attention_mask=torch.ones(1, len(tokens), device="cuda", dtype=torch.long),
|
||||
output_hidden_states=True,
|
||||
use_cache=False,
|
||||
logits_to_keep=1,
|
||||
)
|
||||
if edit:
|
||||
inputs.update(
|
||||
pixel_values=torch.randn(16, 12, device="cuda"),
|
||||
image_grid_thw=torch.tensor([[1, 4, 4]], device="cuda"),
|
||||
)
|
||||
with set_forward_context(
|
||||
current_timestep=None, attn_metadata=None, forward_batch=Req(prompt="test")
|
||||
):
|
||||
expected = reference(**inputs).hidden_states[-1]
|
||||
actual = model(**inputs).hidden_states[-1]
|
||||
torch.testing.assert_close(actual, expected, atol=2e-5, rtol=2e-5)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("height", [4, 5])
|
||||
@pytest.mark.parametrize("residual", [False, True])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float64])
|
||||
@torch.no_grad()
|
||||
def test_vae_spatial_shard_matches_full_decode(height, residual, dtype):
|
||||
arch = QwenImage21VAEArchConfig(
|
||||
base_dim=4,
|
||||
decoder_base_dim=4,
|
||||
z_dim=4,
|
||||
dim_mult=(1, 2, 4, 4, 4),
|
||||
num_res_blocks=1,
|
||||
temperal_downsample=(False, True, True, True),
|
||||
is_residual=residual,
|
||||
)
|
||||
torch.manual_seed(42)
|
||||
reference = (
|
||||
AutoencoderKLQwenImage21(
|
||||
QwenImage21VAEConfig(arch_config=arch, load_encoder=False)
|
||||
)
|
||||
.cuda()
|
||||
.eval()
|
||||
)
|
||||
parallel = (
|
||||
AutoencoderKLQwenImage21(
|
||||
QwenImage21VAEConfig(
|
||||
arch_config=arch,
|
||||
load_encoder=False,
|
||||
parallel_decode_mode="spatial_shard",
|
||||
)
|
||||
)
|
||||
.cuda()
|
||||
.eval()
|
||||
)
|
||||
parallel.load_state_dict(reference.state_dict())
|
||||
assert parallel.spatial_parallel
|
||||
z = torch.randn(1, 4, 1, height, 4, device="cuda")
|
||||
torch.distributed.broadcast(z, src=0)
|
||||
expected = reference.to(dtype).decode(z.to(dtype))
|
||||
actual = parallel.to(dtype).decode(z.to(dtype))
|
||||
assert actual.shape == expected.shape == (1, 4, 1, height * 16, 64)
|
||||
# full and sharded convolutions select different FP32 reduction kernels
|
||||
tolerance = 1e-10 if dtype == torch.float64 else 1e-4
|
||||
torch.testing.assert_close(actual, expected, atol=tolerance, rtol=tolerance)
|
||||
@@ -0,0 +1,91 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.diffusion import (
|
||||
BitExactFusionGate,
|
||||
can_use_fused_complex_rope,
|
||||
fused_complex_rope,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.models.dits import qwen_image21
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.version.hip is not None,
|
||||
reason="NVIDIA CUDA required",
|
||||
)
|
||||
|
||||
|
||||
def reference(x, rope):
|
||||
z = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
||||
return torch.view_as_real(z * rope[None, :, None]).flatten(-2).to(x.dtype)
|
||||
|
||||
|
||||
def inputs(shape, dtype):
|
||||
torch.manual_seed(42)
|
||||
x = torch.randn(shape, device="cuda", dtype=dtype)
|
||||
# a contiguous slice retains the nonzero cache offset used by SP ranks
|
||||
angles = torch.randn(shape[1] + 5, shape[-1] // 2, device="cuda") * 20
|
||||
rope = torch.polar(torch.ones_like(angles), angles)[5:]
|
||||
return x, rope
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32])
|
||||
@pytest.mark.parametrize(
|
||||
"shape", [(1, 1, 1, 32), (2, 17, 3, 64), (1, 257, 16, 128), (1, 4096, 32, 128)]
|
||||
)
|
||||
def test_complex_rope_matches_complex_multiplication(dtype, shape):
|
||||
x, rope = inputs(shape, dtype)
|
||||
assert can_use_fused_complex_rope(x, rope)
|
||||
actual = fused_complex_rope(x, rope)
|
||||
torch.testing.assert_close(actual, reference(x, rope), atol=0, rtol=0)
|
||||
|
||||
|
||||
def test_complex_rope_layout_guards():
|
||||
x, rope = inputs((2, 17, 3, 64), torch.bfloat16)
|
||||
assert not can_use_fused_complex_rope(x.cpu(), rope.cpu())
|
||||
assert not can_use_fused_complex_rope(x.double(), rope)
|
||||
assert not can_use_fused_complex_rope(x, rope.to(torch.complex128))
|
||||
assert not can_use_fused_complex_rope(x[:, ::2], rope[::2])
|
||||
assert not can_use_fused_complex_rope(x, rope[:-1])
|
||||
assert not can_use_fused_complex_rope(x[:, :0], rope[:0])
|
||||
|
||||
|
||||
def test_complex_rope_compile_and_graph_replay():
|
||||
x, rope = inputs((1, 257, 8, 128), torch.bfloat16)
|
||||
compiled = torch.compile(fused_complex_rope, fullgraph=True)
|
||||
torch.testing.assert_close(compiled(x, rope), reference(x, rope), atol=0, rtol=0)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
out = fused_complex_rope(x, rope)
|
||||
x.normal_()
|
||||
graph.replay()
|
||||
torch.testing.assert_close(out, reference(x, rope), atol=0, rtol=0)
|
||||
|
||||
|
||||
def test_qwen21_rope_first_sight_verification(monkeypatch):
|
||||
x, rope = inputs((1, 257, 8, 128), torch.bfloat16)
|
||||
gate = BitExactFusionGate("test complex RoPE")
|
||||
monkeypatch.setattr(qwen_image21, "_ROPE_FUSION", gate)
|
||||
torch.testing.assert_close(
|
||||
qwen_image21.apply_rope(x, rope), reference(x, rope), atol=0, rtol=0
|
||||
)
|
||||
assert gate.verified and not gate.disabled
|
||||
|
||||
gate = BitExactFusionGate("test mismatched RoPE")
|
||||
monkeypatch.setattr(qwen_image21, "_ROPE_FUSION", gate)
|
||||
monkeypatch.setattr(
|
||||
qwen_image21, "fused_complex_rope", lambda x, rope: torch.zeros_like(x)
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
qwen_image21.apply_rope(x, rope), reference(x, rope), atol=0, rtol=0
|
||||
)
|
||||
assert gate.disabled and not gate.verified
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -0,0 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
|
||||
from sglang.kernels.ops.diffusion import (
|
||||
can_use_fused_layernorm_modulate,
|
||||
fused_layernorm_modulate,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.version.hip is not None,
|
||||
reason="NVIDIA CUDA required",
|
||||
)
|
||||
|
||||
|
||||
def reference(x, scale, shift, eps):
|
||||
out = F.layer_norm(x, (x.shape[-1],), eps=eps) * (1 + scale[:, None])
|
||||
return out if shift is None else out + shift[:, None]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("shape", [(1, 1, 128), (1, 4359, 4096), (2, 1024, 4096)])
|
||||
@pytest.mark.parametrize("amplitude,eps", [(1e-4, 1e-6), (1.0, 1e-6), (100.0, 1e-5)])
|
||||
@pytest.mark.parametrize("has_shift", [False, True])
|
||||
def test_modulation_preserves_bits(shape, amplitude, eps, has_shift):
|
||||
torch.manual_seed(42)
|
||||
x = torch.randn(shape, device="cuda", dtype=torch.bfloat16) * amplitude
|
||||
modulation = torch.randn(shape[0], 4 * shape[-1], device="cuda", dtype=x.dtype)
|
||||
scale, shift = modulation.chunk(4, dim=-1)[:2]
|
||||
scale[:, :3] = torch.tensor([-1, 0, 1], device=x.device, dtype=x.dtype)
|
||||
if not has_shift:
|
||||
shift = None
|
||||
assert can_use_fused_layernorm_modulate(x, scale, shift)
|
||||
actual = fused_layernorm_modulate(x, scale, shift, eps)
|
||||
expected = reference(x, scale, shift, eps)
|
||||
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
|
||||
|
||||
|
||||
def test_scale_only_preserves_signed_zero():
|
||||
x = torch.ones(1, 17, 128, device="cuda", dtype=torch.bfloat16)
|
||||
scale = torch.full((1, 128), -2, device=x.device, dtype=x.dtype)
|
||||
actual = fused_layernorm_modulate(x, scale, None, 1e-6)
|
||||
expected = reference(x, scale, None, 1e-6)
|
||||
assert torch.signbit(expected).all()
|
||||
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
|
||||
|
||||
|
||||
def test_scale_only_layout_guards():
|
||||
x = torch.randn(2, 17, 128, device="cuda", dtype=torch.bfloat16)
|
||||
scale = torch.randn(2, 128, device=x.device, dtype=x.dtype)
|
||||
assert can_use_fused_layernorm_modulate(x, scale, None)
|
||||
assert not can_use_fused_layernorm_modulate(x.cpu(), scale.cpu(), None)
|
||||
assert not can_use_fused_layernorm_modulate(x.float(), scale.float(), None)
|
||||
assert not can_use_fused_layernorm_modulate(x[:, ::2], scale, None)
|
||||
assert not can_use_fused_layernorm_modulate(x, scale.float(), None)
|
||||
assert not can_use_fused_layernorm_modulate(x, scale[:, :-1], None)
|
||||
assert not can_use_fused_layernorm_modulate(x[:, :0], scale, None)
|
||||
assert not can_use_fused_layernorm_modulate(x, scale, scale.float())
|
||||
strided = torch.empty(2, 256, device=x.device, dtype=x.dtype)[:, :128]
|
||||
assert not can_use_fused_layernorm_modulate(x, scale, strided)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("has_shift", [False, True])
|
||||
def test_compile_and_graph_replay(has_shift):
|
||||
x = torch.randn(2, 17, 128, device="cuda", dtype=torch.bfloat16)
|
||||
modulation = torch.randn(2, 512, device=x.device, dtype=x.dtype)
|
||||
scale, shift = modulation.chunk(4, dim=-1)[:2]
|
||||
if not has_shift:
|
||||
shift = None
|
||||
compiled = torch.compile(fused_layernorm_modulate, fullgraph=True)
|
||||
expected = reference(x, scale, shift, 1e-6)
|
||||
actual = compiled(x, scale, shift, 1e-6)
|
||||
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
out = fused_layernorm_modulate(x, scale, shift, 1e-6)
|
||||
x.normal_()
|
||||
modulation.normal_()
|
||||
graph.replay()
|
||||
expected = reference(x, scale, shift, 1e-6)
|
||||
assert torch.equal(out.view(torch.int16), expected.view(torch.int16))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -37,8 +37,10 @@ import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image
|
||||
import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat_image
|
||||
import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module
|
||||
import sglang.multimodal_gen.runtime.models.dits.qwen_image as qwen_image
|
||||
import sglang.multimodal_gen.runtime.models.dits.qwen_image21 as qwen_image21
|
||||
import sglang.multimodal_gen.runtime.models.dits.sana as sana
|
||||
from sglang.kernels.ops.diffusion import (
|
||||
BitExactFusionGate,
|
||||
can_use_fused_layernorm_modulate,
|
||||
can_use_fused_qk_head_layernorm,
|
||||
can_use_fused_rmsnorm_scale_shift,
|
||||
@@ -1479,5 +1481,97 @@ def test_autoencoder_kl_fastpath_install():
|
||||
assert torch.equal(opt.decode(z), ref)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_qwen21_qk_norm_verifies_and_preserves_native_fallback(monkeypatch):
|
||||
x = torch.randn(1, 257, 8, 128, device="cuda", dtype=torch.bfloat16)
|
||||
norm = qwen_image21.RMSNorm(
|
||||
128, 1e-6, cast_x_before_out_mul=True, force_native=True
|
||||
).to(device=x.device, dtype=x.dtype)
|
||||
norm.weight.normal_()
|
||||
expected = norm(x)
|
||||
gate = BitExactFusionGate("test Q/K norm")
|
||||
monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate)
|
||||
assert torch.equal(qwen_image21.apply_qk_norm(x, norm), expected)
|
||||
assert gate.verified and not gate.disabled
|
||||
x.normal_()
|
||||
assert torch.equal(qwen_image21.apply_qk_norm(x, norm), norm(x))
|
||||
|
||||
gate = BitExactFusionGate("test mismatched Q/K norm")
|
||||
monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate)
|
||||
monkeypatch.setattr(
|
||||
qwen_image21,
|
||||
"rmsnorm_preserve_reduction",
|
||||
lambda x, weight, eps: torch.zeros_like(x),
|
||||
)
|
||||
assert torch.equal(qwen_image21.apply_qk_norm(x, norm), norm(x))
|
||||
assert gate.disabled and not gate.verified
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_qwen21_qk_norm_does_not_verify_during_capture(monkeypatch):
|
||||
x = torch.randn(1, 17, 2, 128, device="cuda", dtype=torch.bfloat16)
|
||||
norm = qwen_image21.RMSNorm(
|
||||
128, 1e-6, cast_x_before_out_mul=True, force_native=True
|
||||
).to(device=x.device, dtype=x.dtype)
|
||||
norm(x)
|
||||
gate = BitExactFusionGate("test captured Q/K norm")
|
||||
monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
out = qwen_image21.apply_qk_norm(x, norm)
|
||||
assert not gate.verified and not gate.disabled
|
||||
x.normal_()
|
||||
graph.replay()
|
||||
assert torch.equal(out, norm(x))
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_qwen21_modulation_verifies_and_preserves_native_fallback(monkeypatch):
|
||||
x = torch.randn(1, 257, 4096, device="cuda", dtype=torch.bfloat16)
|
||||
scale = torch.randn(1, 1, 4096, device=x.device, dtype=x.dtype)
|
||||
norm = torch.nn.LayerNorm(4096, eps=1e-6, elementwise_affine=False).cuda()
|
||||
gate = BitExactFusionGate("test scale-only modulation")
|
||||
monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate)
|
||||
expected = norm(x) * (1 + scale)
|
||||
actual = qwen_image21.apply_modulation(x, norm, scale)
|
||||
assert torch.equal(actual.view(torch.int16), expected.view(torch.int16))
|
||||
assert gate.verified and not gate.disabled
|
||||
x.normal_()
|
||||
scale.normal_()
|
||||
assert torch.equal(
|
||||
qwen_image21.apply_modulation(x, norm, scale), norm(x) * (1 + scale)
|
||||
)
|
||||
|
||||
gate = BitExactFusionGate("test mismatched modulation")
|
||||
monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate)
|
||||
monkeypatch.setattr(
|
||||
qwen_image21,
|
||||
"fused_layernorm_modulate",
|
||||
lambda x, scale, shift, eps: torch.zeros_like(x),
|
||||
)
|
||||
assert torch.equal(
|
||||
qwen_image21.apply_modulation(x, norm, scale), norm(x) * (1 + scale)
|
||||
)
|
||||
assert gate.disabled and not gate.verified
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_qwen21_modulation_does_not_verify_during_capture(monkeypatch):
|
||||
x = torch.randn(1, 17, 128, device="cuda", dtype=torch.bfloat16)
|
||||
scale = torch.randn(1, 1, 128, device=x.device, dtype=x.dtype)
|
||||
norm = torch.nn.LayerNorm(128, eps=1e-6, elementwise_affine=False).cuda()
|
||||
norm(x)
|
||||
gate = BitExactFusionGate("test captured modulation")
|
||||
monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
out = qwen_image21.apply_modulation(x, norm, scale)
|
||||
assert not gate.verified and not gate.disabled
|
||||
x.normal_()
|
||||
scale.normal_()
|
||||
graph.replay()
|
||||
assert torch.equal(out, norm(x) * (1 + scale))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.diffusion import (
|
||||
can_use_rmsnorm_preserve_reduction,
|
||||
rmsnorm_preserve_reduction,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available() or torch.version.hip is not None,
|
||||
reason="NVIDIA CUDA required",
|
||||
)
|
||||
|
||||
|
||||
def reference(x, weight, eps):
|
||||
value = x.float()
|
||||
variance = value.pow(2).mean(dim=-1, keepdim=True)
|
||||
return weight * (value * torch.rsqrt(variance + eps)).to(x.dtype)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
|
||||
@pytest.mark.parametrize("shape", [(1, 128), (2, 17, 3, 128), (1, 4096, 32, 128)])
|
||||
@pytest.mark.parametrize("scale,eps", [(1e-4, 1e-6), (1.0, 1e-6), (100.0, 1e-5)])
|
||||
def test_preserves_native_reduction_and_rounding(dtype, shape, scale, eps):
|
||||
torch.manual_seed(42)
|
||||
x = (torch.randn(shape, device="cuda") * scale).to(dtype)
|
||||
weight = torch.randn(shape[-1], device="cuda", dtype=dtype)
|
||||
assert can_use_rmsnorm_preserve_reduction(x, weight)
|
||||
actual = rmsnorm_preserve_reduction(x, weight, eps)
|
||||
torch.testing.assert_close(actual, reference(x, weight, eps), atol=0, rtol=0)
|
||||
|
||||
|
||||
def test_layout_guards_and_offset():
|
||||
x = torch.randn(259, 128, device="cuda", dtype=torch.bfloat16)[2:]
|
||||
weight = torch.randn(128, device="cuda", dtype=x.dtype)
|
||||
assert can_use_rmsnorm_preserve_reduction(x, weight)
|
||||
torch.testing.assert_close(
|
||||
rmsnorm_preserve_reduction(x, weight, 1e-6),
|
||||
reference(x, weight, 1e-6),
|
||||
atol=0,
|
||||
rtol=0,
|
||||
)
|
||||
assert not can_use_rmsnorm_preserve_reduction(x.cpu(), weight.cpu())
|
||||
assert not can_use_rmsnorm_preserve_reduction(x.float(), weight.float())
|
||||
assert not can_use_rmsnorm_preserve_reduction(x[:, ::2], weight[::2])
|
||||
assert not can_use_rmsnorm_preserve_reduction(x, weight.float())
|
||||
assert not can_use_rmsnorm_preserve_reduction(x, weight[:-1])
|
||||
assert not can_use_rmsnorm_preserve_reduction(x[:0], weight)
|
||||
|
||||
|
||||
def test_compile_and_graph_replay():
|
||||
x = torch.randn(257, 128, device="cuda", dtype=torch.bfloat16)
|
||||
weight = torch.randn(128, device="cuda", dtype=x.dtype)
|
||||
compiled = torch.compile(rmsnorm_preserve_reduction, fullgraph=True)
|
||||
torch.testing.assert_close(
|
||||
compiled(x, weight, 1e-6), reference(x, weight, 1e-6), atol=0, rtol=0
|
||||
)
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
out = rmsnorm_preserve_reduction(x, weight, 1e-6)
|
||||
x.normal_()
|
||||
weight.normal_()
|
||||
graph.replay()
|
||||
torch.testing.assert_close(out, reference(x, weight, 1e-6), atol=0, rtol=0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user