diff --git a/docs/cookbook/diffusion/CircleStone/Anima.mdx b/docs/cookbook/diffusion/CircleStone/Anima.mdx new file mode 100644 index 000000000000..b32b48390eaa --- /dev/null +++ b/docs/cookbook/diffusion/CircleStone/Anima.mdx @@ -0,0 +1,139 @@ +--- +title: Anima +description: "Deploy Anima Base v1.0 with SGLang Diffusion for anime and illustration generation, using native components and single- or multi-GPU execution." +--- + +import { DiffusionModelTags } from '/src/snippets/diffusion/model-tags.jsx'; +import { Deployment } from '/src/snippets/_deployment.jsx'; +import { config } from '/src/snippets/configs/CircleStone/anima.jsx'; + + + +## 1. Quick start + +Install the diffusion dependencies on Linux with NVIDIA CUDA, then install this +integration from its source checkout: + +```bash Installation +uv pip install "sglang[diffusion]" --prerelease=allow +uv pip install -e "python[diffusion]" +``` + +Use the official `circlestone-labs/Anima-Base-v1.0-Diffusers` checkpoint. SGLang reads +its `modular_model_index.json` directly; no checkpoint conversion or newer Diffusers +runtime is required. + +Start with **Resident + Compiled** for repeated generation on RTX 5090, DGX Spark, +or H200. Select **Eager** for shorter startup or frequent shape changes. Spark's +128 GB is unified memory shared with the CPU, not dedicated VRAM. See the +[measurements and tradeoffs](#5-measured-tuning) below. + + + +## 2. Model capabilities + +[Anima](https://huggingface.co/circlestone-labs/Anima) is CircleStone Labs and Comfy +Org's text-to-image model for anime and illustration. It combines a Cosmos Predict2 +transformer with Qwen3 text encoding, a learned T5-token conditioner, and a Qwen-Image +VAE. Prompts can combine tags with natural-language descriptions. + +This integration targets **Base v1.0 text-to-image**. It does not claim support for +the separate Aesthetic or Turbo checkpoints, single-file ComfyUI weights, or image +editing. The checkpoint uses the CircleStone Labs Non-Commercial License; review +the [official license](https://huggingface.co/circlestone-labs/Anima/blob/main/LICENSE.md) +before deployment. + +## 3. Sampling + +Defaults are 1024 x 1024, 30 steps, guidance scale 4, and an empty negative prompt. +Width and height should be divisible by 16. `max_sequence_length` defaults to 512 +and accepts 1 through 4096. The text conditioner pads short sequences to 512 tokens; +these padding positions remain part of the transformer's cross-attention, matching +the official implementation. + +For offline generation: + +```bash Generate +sglang generate \ + --model-path circlestone-labs/Anima-Base-v1.0-Diffusers \ + --prompt "masterpiece, best quality, safe, watercolor landscape, a quiet seaside village at sunset" \ + --seed 42 \ + --save-output +``` + +Use the same seed, generator device, dimensions, scheduler settings, and precision +when comparing runtimes. Latents and scheduler updates remain FP32; the transformer, +text components, and VAE default to BF16. Different attention kernels can produce +small floating-point differences that accumulate during denoising. + +## 4. Runtime features + +The pipeline reuses SGLang's native Qwen3 encoder, Qwen-Image VAE, component loaders, +denoising loop, and residency management. Anima's transformer and text conditioner +are native modules, not wrappers around Diffusers models. + +For multi-GPU execution, choose TP to shard transformer weights or Ulysses/Ring to +shard image tokens. The transformer has 16 heads, so `tp_size * ulysses_degree` +must divide 16. CFG parallelism additionally splits conditional and unconditional +denoising and requires guidance greater than 1. The total GPU count must match the +selected parallel topology. + +CPU and layerwise offload trade memory for transfers. The additional +`text_conditioner` component accepts the same residency controls as other native +components. Cache-DiT and quantized attention are approximate optimizations; assess +image quality for your prompts before enabling them. See +[performance optimization](/docs/sglang-diffusion/performance-optimization) for the +shared controls. H200 functional checks cover TP, Ulysses, Ring, CFG parallelism, +encoder folding, parallel tiled VAE decode, layerwise offload, FlashAttention, +Torch SDPA, SageAttention, Cache-DiT, and breakable CUDA graphs. These checks are +not a quality guarantee for approximate optimizations. Combined TP and SP, +multi-node execution, non-NVIDIA devices, and third-party LoRAs remain unverified. + +Breakable CUDA graphs reuse the conditioner's actual sequence length, without +additional text-bucket padding. Requests with uncaptured shapes fall back to eager +execution. Standard short prompts use the same 512-token conditioning shape. + +## 5. Measured tuning + +For repeated requests, compilation was faster than eager execution on all three +tested platforms. The picker keeps eager available because initial compilation +and recompilation for new shapes can take minutes and use additional CPU memory. +No CPU offload, quantized attention, or Cache-DiT is enabled in the recommended +recipes. + +The following are warm, sequential HTTP request medians: 1024 x 1024, 30 steps, +CFG 4, one image, CPU generator, and PNG/base64 output. Startup is excluded. + +| GPU | Recommended recipe | Original eager | Recommended | Peak allocated | +| --- | --- | --- | --- | --- | +| RTX 5090, 32 GB | cuDNN SDPA + compile, tiled VAE | 6.99 s | 6.10 s | 5.7 GiB | +| DGX Spark, 128 GB unified | Torch SDPA + compile, tiled VAE | 26.03 s | 17.41 s | 5.6 GiB | +| H200, 141 GB | FlashAttention + compile, untiled VAE | 3.22 s | 2.22 s | 10.2 GiB | + +Measurements used five timed requests after warmup, PyTorch +2.13.0+cu130, and the native Anima implementation at `214c9a48bb1` (Spark used +`7f3f903f039`, with identical runtime code). Results are workload-specific, not +throughput-at-saturation measurements. The original eager baseline uses Torch +SDPA on 5090/Spark and FlashAttention on H200, with VAE tiling enabled. Peak +allocated memory excludes the CUDA context, reserved pool, and compilation CPU +memory; it is not total device usage. + +On **two NVLink-connected H200s**, Auto selects CFG parallelism instead of TP: +the compiled, untiled recipe measured **1.28 s**. Untiled decoding trades about +4-5 GiB more allocated memory for lower latency and is verified at 1024 x 1024, +one output. The tiled recipes additionally passed repeated 512 x 512 and +1536 x 1536 requests, plus two outputs at 1024 x 1024. Use **Tiled** for that +verified scope. At 512 x 512, two-GPU CFG eager was faster than compiled; +more GPUs or compilation are not universally better. + +The recommended recipes change floating-point kernels and, on H200, VAE tiling. +They are **not pixel-identical to eager**: across three fixed prompt/seed pairs, +the single-GPU recipes above measured PSNR 21.95-35.44 dB and SSIM 0.851-0.969 +against their same-device eager outputs. +These measure output differences, not perceptual quality guarantees. Use eager +with the same backend and tiling settings when reproducing an eager reference. + +Breakable CUDA graphs were near parity at 1024 x 1024 while consuming about +1.2 GiB more allocated GPU memory for a single captured shape. They are not the +default. Cache-DiT can accelerate further, but changes the denoising computation; +validate it separately against your quality requirements. diff --git a/docs/docs.json b/docs/docs.json index b57f7b81f68c..37c3163a7dcd 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -1521,6 +1521,12 @@ "cookbook/diffusion/Qwen-Image/Qwen-Image-Edit" ] }, + { + "group": "CircleStone", + "pages": [ + "cookbook/diffusion/CircleStone/Anima" + ] + }, { "group": "SenseNova", "tag": "NEW", diff --git a/docs/src/snippets/_deployment.jsx b/docs/src/snippets/_deployment.jsx index 9ec2a3c61a5a..a70ed97ba96b 100644 --- a/docs/src/snippets/_deployment.jsx +++ b/docs/src/snippets/_deployment.jsx @@ -1865,16 +1865,19 @@ export const Deployment = ({ config, benchmarks }) => { && Number(sel.nodes) === recommendedRecipe.nodes && Number(sel.gpus_per_node) === recommendedRecipe.gpus_per_node && sel.topology_mode === "auto" - && ["auto", recommendedRecipe.placement].includes(sel.placement) - && sel.attention === "platform" - && sel.precision === "native" - && ["auto", recommendedRecipe.encoder].includes(sel.encoder) - && sel.execution === "eager"; + && serveDims.every((dim) => { + const expected = recommendedRecipe[dim.id] ?? dim.default; + if (dim.id === "attention") return sel.attention === "platform"; + return expected === undefined || sel[dim.id] === expected + || (["placement", "encoder"].includes(dim.id) && sel[dim.id] === "auto"); + }); const restoreRecommendedRecipe = () => { if (!recommendedRecipe) return; setSel((prev) => reseatHiddenPicks(normalizeBuilderSelection({ ...prev, + ...Object.fromEntries(serveDims + .map((dim) => [dim.id, recommendedRecipe[dim.id] ?? dim.default ?? dim.options?.[0]?.id])), nodes: recommendedRecipe.nodes, gpus_per_node: recommendedRecipe.gpus_per_node, topology_mode: "auto", @@ -1885,7 +1888,7 @@ export const Deployment = ({ config, benchmarks }) => { attention: "platform", precision: "native", encoder: recommendedRecipe.encoder || "auto", - execution: "eager", + execution: recommendedRecipe.execution || "eager", }))); }; diff --git a/docs/src/snippets/configs/CircleStone/anima.jsx b/docs/src/snippets/configs/CircleStone/anima.jsx new file mode 100644 index 000000000000..601bdae1c1ef --- /dev/null +++ b/docs/src/snippets/configs/CircleStone/anima.jsx @@ -0,0 +1,141 @@ +export const config = { + modelName: "Anima Base v1.0", + supportedHardware: ["rtx5090", "dgx-spark", "h200"], + hardware: [ + { id: "rtx5090", label: "RTX 5090", vram: "32GB", vendor: "consumer" }, + { id: "dgx-spark", label: "DGX Spark", vram: "128GB unified", vendor: "blackwell" }, + ], + groupHardware: false, + matchDims: [], + overlayDims: [ + { + id: "weights", title: "Checkpoint", scope: "base", default: "default", + description: "Official self-contained Diffusers checkpoint.", + options: [{ id: "default", label: "Base v1.0" }], + }, + { + id: "mode", title: "Request mode", scope: "base", default: "text", + options: [{ id: "text", label: "Text to image" }], + }, + { + id: "placement", title: "Placement", scope: "serve", default: "resident", + description: "Keep the pipeline resident when it fits. Use offload for a smaller memory budget.", + learnMore: "#4-runtime-features", + options: [ + { id: "resident", label: "Resident", flags: ["--performance-mode speed"], recommended: true }, + { id: "offload", label: "Layerwise offload", flags: ["--performance-mode memory", "--dit-layerwise-offload true"], soft: true, softReason: "This exact HTTP recipe is unverified." }, + ], + }, + { + id: "execution", title: "Execution", scope: "serve", default: "compile", + description: "Compile for repeated requests. Eager avoids compilation at startup and for new shapes.", + learnMore: "#5-measured-tuning", + options: [ + { id: "compile", label: "Compiled", flags: ["--enable-torch-compile true"], recommended: true }, + { id: "eager", label: "Eager", flags: [] }, + ], + }, + { + id: "attention", title: "Attention", scope: "serve", default: "platform", + description: "Exact attention backends can differ in floating-point reduction order.", + options: [ + { id: "platform", label: "Platform default", flags: (s) => [`--attention-backend ${s.hw === "h200" ? "fa" : s.hw === "rtx5090" ? "torch_cudnn_sdpa" : "torch_sdpa"}`], recommended: true }, + { id: "fa", label: "FlashAttention", flags: ["--attention-backend fa"], soft: (s) => s.hw !== "h200", softReason: "On RTX 5090 and Spark, this selector falls back to Torch SDPA. Select Torch SDPA explicitly." }, + { id: "sdpa", label: "Torch SDPA", flags: ["--attention-backend torch_sdpa"] }, + { id: "cudnn", label: "cuDNN SDPA", flags: ["--attention-backend torch_cudnn_sdpa"] }, + ], + }, + { + id: "cfg", title: "CFG parallelism", scope: "serve", default: "auto", + description: "Auto uses CFG parallelism on two H200 GPUs. Manual TP/SP overrides disable automatic CFG splitting.", + learnMore: "#5-measured-tuning", + options: [ + { id: "auto", label: "Auto", flags: [] }, + { id: "off", label: "Off", flags: [] }, + { id: "on", label: "On", flags: [], disabled: (s) => Number(s.gpus_per_node) % 2 !== 0, disableReason: "CFG parallelism requires an even number of GPUs and guidance greater than 1." }, + ], + }, + { + id: "vae", title: "VAE decoding", scope: "serve", default: "auto", + description: "Auto disables tiling for compiled H200 recipes. Untiled decoding uses about 4-5 GiB more memory at 1024 x 1024.", + learnMore: "#5-measured-tuning", + options: [ + { id: "auto", label: "Auto", flags: [] }, + { id: "tiled", label: "Tiled", flags: [] }, + { id: "full", label: "Untiled", flags: [] }, + ], + }, + { + id: "outputs", title: "Outputs", scope: "request", kind: "number", + default: 1, min: 1, max: 4, options: [], + }, + ], + commandBuilder: { + defaultSelection: { hw: "rtx5090", 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: 8 } }, + verifiedRecipes: [ + { id: "rtx5090-1gpu-resident-cudnn", hw: "rtx5090", nodes: 1, gpus_per_node: 1, tp_size: 1, ulysses_degree: 1, ring_degree: 1, placement: "resident", attention: "cudnn", execution: "compile", cfg: "auto", vae: "auto" }, + { id: "dgx-spark-1gpu-resident-sdpa", hw: "dgx-spark", nodes: 1, gpus_per_node: 1, tp_size: 1, ulysses_degree: 1, ring_degree: 1, placement: "resident", attention: "sdpa", execution: "compile", cfg: "auto", vae: "auto" }, + { id: "h200-1gpu-resident-fa", hw: "h200", nodes: 1, gpus_per_node: 1, tp_size: 1, ulysses_degree: 1, ring_degree: 1, placement: "resident", attention: "fa", execution: "compile", cfg: "auto", vae: "auto" }, + { id: "h200-2gpu-resident-cfg", hw: "h200", nodes: 1, gpus_per_node: 2, tp_size: 1, ulysses_degree: 1, ring_degree: 1, placement: "resident", attention: "fa", execution: "compile", cfg: "auto", vae: "auto" }, + ], + cfgDegree: (s) => s.cfg === "on" || (s.cfg === "auto" && s.topology_mode !== "manual" && s.hw === "h200" && Number(s.gpus_per_node) === 2) ? 2 : 1, + vaeTiling: (s) => s.vae === "tiled" || (s.vae === "auto" && (s.hw !== "h200" || s.execution !== "compile")), + autoTopology: (s) => ({ tp_size: Number(s.gpus_per_node) / config.commandBuilder.resource.cfgDegree(s), ulysses_degree: 1, ring_degree: 1 }), + validateTopology: (s, t) => { + const errors = []; + if (Number(s.nodes) !== 1) errors.push("This picker covers single-node deployment only."); + if (s.hw === "dgx-spark" && Number(s.gpus_per_node) !== 1) errors.push("DGX Spark has one GPU per node."); + const cfg = config.commandBuilder.resource.cfgDegree(s); + if (!Number.isInteger(t.tp_size) || t.tp_size < 1 || Number(s.nodes) * Number(s.gpus_per_node) !== cfg * t.tp_size * t.ulysses_degree * t.ring_degree) errors.push("GPU count must equal CFG * TP * Ulysses * Ring."); + if (16 % (t.tp_size * t.ulysses_degree)) errors.push("TP * Ulysses must divide Anima's 16 attention heads."); + if (t.ring_degree > 1 && (s.hw !== "h200" || ["sdpa", "cudnn"].includes(s.attention))) errors.push("Ring requires FlashAttention."); + return errors; + }, + }, + resolveDeployment: (s) => { + const r = config.commandBuilder.resource; + const attention = s.attention === "platform" ? (s.hw === "h200" ? "fa" : s.hw === "rtx5090" ? "cudnn" : "sdpa") : s.attention; + const t = s.topology_mode === "manual" + ? { tp_size: Number(s.tp_size), ulysses_degree: Number(s.ulysses_degree), ring_degree: Number(s.ring_degree) } + : r.autoTopology(s); + const errors = r.validateTopology(s, t); + const cfg = r.cfgDegree(s); + const tiled = r.vaeTiling(s); + const world = Number(s.nodes) * Number(s.gpus_per_node); + const recipe = r.verifiedRecipes.find((v) => v.hw === s.hw && v.gpus_per_node === Number(s.gpus_per_node) && v.tp_size === t.tp_size && v.ulysses_degree === t.ulysses_degree && v.ring_degree === t.ring_degree && v.placement === s.placement && r.cfgDegree(v) === cfg); + const attentions = world > 1 ? ["fa"] : s.execution === "eager" + ? ["sdpa", "cudnn", ...(s.hw === "h200" ? ["fa"] : [])] + : s.hw === "rtx5090" ? ["sdpa", "cudnn"] : [s.hw === "h200" ? "fa" : "sdpa"]; + const multiOutput = tiled && (s.execution === "compile" || attention === "fa" || (attention === "sdpa" && s.hw !== "h200")); + const vaeVerified = tiled || (s.execution === "compile" && attention === (s.hw === "h200" ? "fa" : "sdpa")); + const verified = !!recipe && attentions.includes(attention) && vaeVerified && Number(s.outputs) <= (multiOutput ? 2 : 1) && errors.length === 0; + const flags = ["--model-path {{MODEL_NAME}}"]; + if (world > 1) flags.push(`--num-gpus ${world}`, `--tp-size ${t.tp_size}`, `--ulysses-degree ${t.ulysses_degree}`, `--ring-degree ${t.ring_degree}`); + if (cfg > 1) flags.push("--enable-cfg-parallel"); + if (!tiled) flags.push("--vae-tiling false", "--vae-sp false"); + flags.push("--host {{HOST_IP}}", "--port {{PORT}}"); + return { + match: { hw: s.hw }, nnodes: Number(s.nodes), verified, flags, + builder: { + topology: t, topologySummary: `CFG ${cfg} / TP ${t.tp_size} / Ulysses ${t.ulysses_degree} / Ring ${t.ring_degree}`, + errors, warnings: verified ? [] : ["This exact HTTP recipe has not been verified."], + verification: { serve: verified ? "verified" : "unverified", request: verified ? "verified" : "unverified" }, + resolvedSettings: { attention: s.attention === "platform" ? ({ fa: "FlashAttention", sdpa: "Torch SDPA", cudnn: "cuDNN SDPA" }[attention]) : undefined, cfg: cfg > 1 ? "2-way CFG" : "Off", vae: tiled ? "Tiled" : "Untiled" }, + }, + }; + }, + }, + modelNames: { default: "circlestone-labs/Anima-Base-v1.0-Diffusers" }, + placeholders: { + HOST_IP: { target: "command", label: "Bind host", default: "0.0.0.0" }, + PORT: { target: "command", label: "Bind port", default: "30000" }, + CURL_HOST: { target: "curl", label: "Server host", default: "localhost" }, + CURL_PORT: { target: "curl", label: "Server port", default: "30000" }, + }, + curl: (s) => `curl -sS http://{{CURL_HOST}}:{{CURL_PORT}}/v1/images/generations \\ + -H 'Content-Type: application/json' \\ + -d '${JSON.stringify({ model: "{{MODEL_NAME}}", prompt: "masterpiece, best quality, safe, watercolor landscape, a quiet seaside village at sunset", size: "1024x1024", n: Number(s.outputs), seed: 42, generator_device: "cpu", output_format: "png", response_format: "b64_json" }, null, 2)}'`, + cells: [], +}; diff --git a/docs/src/snippets/diffusion/model-catalog.jsx b/docs/src/snippets/diffusion/model-catalog.jsx index 5c34eb710c0d..5ed174b4b569 100644 --- a/docs/src/snippets/diffusion/model-catalog.jsx +++ b/docs/src/snippets/diffusion/model-catalog.jsx @@ -3,6 +3,11 @@ export const DiffusionModelCatalog = ({ category }) => { const MODEL_CATALOG = { image: [ + { + name: "Anima", + modelIds: ["circlestone-labs/Anima-Base-v1.0-Diffusers"], + cookbook: "/cookbook/diffusion/CircleStone/Anima", + }, { name: "Ming-Image", modelIds: ["inclusionAI/Ming-Image-0.1-Design", "inclusionAI/Ming-Image-0.1-Design-Layer"], diff --git a/python/sglang/cli/utils.py b/python/sglang/cli/utils.py index 5df252a3ea7a..eebaecea02fe 100644 --- a/python/sglang/cli/utils.py +++ b/python/sglang/cli/utils.py @@ -43,15 +43,13 @@ def _is_diffusion_model_from_registry(model_path: str) -> bool: def _is_diffusers_model_dir(model_dir: str) -> bool: - """Check if a local directory contains a valid diffusers model_index.json.""" - config_path = os.path.join(model_dir, "model_index.json") - if not os.path.exists(config_path): - return False - - with open(config_path) as f: - config = json.load(f) - - return "_diffusers_version" in config + """Check for a standard or modular Diffusers pipeline index.""" + for filename in ("model_index.json", "modular_model_index.json"): + config_path = os.path.join(model_dir, filename) + if os.path.isfile(config_path): + with open(config_path) as f: + return "_diffusers_version" in json.load(f) + return False def _is_gated_diffusion_repo(repo_id: str) -> bool: diff --git a/python/sglang/multimodal_gen/README.md b/python/sglang/multimodal_gen/README.md index 868edc17a98b..e14fbfb197d6 100644 --- a/python/sglang/multimodal_gen/README.md +++ b/python/sglang/multimodal_gen/README.md @@ -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 / 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 + - Broad model support: Wan, FastWan, FLUX, Qwen-Image / Qwen-Image 2.1, LongCat-Image, Z-Image, Anima, 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: diff --git a/python/sglang/multimodal_gen/configs/models/adapter/anima.py b/python/sglang/multimodal_gen/configs/models/adapter/anima.py new file mode 100644 index 000000000000..9e63d680bf9c --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/adapter/anima.py @@ -0,0 +1,28 @@ +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.models.adapter.base import ( + AdapterArchConfig, + AdapterConfig, +) + + +@dataclass +class AnimaTextConditionerArchConfig(AdapterArchConfig): + source_dim: int = 1024 + target_dim: int = 1024 + model_dim: int = 1024 + num_layers: int = 6 + num_attention_heads: int = 16 + mlp_ratio: float = 4.0 + target_vocab_size: int = 32128 + use_self_attention: bool = True + use_layer_norm: bool = False + min_sequence_length: int = 512 + + +@dataclass +class AnimaTextConditionerConfig(AdapterConfig): + arch_config: AnimaTextConditionerArchConfig = field( + default_factory=AnimaTextConditionerArchConfig + ) diff --git a/python/sglang/multimodal_gen/configs/models/dits/anima.py b/python/sglang/multimodal_gen/configs/models/dits/anima.py new file mode 100644 index 000000000000..392b2991f393 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/dits/anima.py @@ -0,0 +1,34 @@ +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field + +from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig + + +@dataclass +class AnimaArchConfig(DiTArchConfig): + in_channels: int = 16 + out_channels: int = 16 + num_attention_heads: int = 16 + attention_head_dim: int = 128 + num_layers: int = 28 + mlp_ratio: float = 4.0 + text_embed_dim: int = 1024 + adaln_lora_dim: int = 256 + patch_size: tuple[int, int, int] = (1, 2, 2) + max_size: tuple[int, int, int] = (128, 240, 240) + rope_scale: tuple[float, float, float] = (1.0, 4.0, 4.0) + concat_padding_mask: bool = True + extra_pos_embed_type: str | None = None + use_crossattn_projection: bool = False + img_context_dim_in: int | None = None + + def __post_init__(self): + super().__post_init__() + self.hidden_size = self.num_attention_heads * self.attention_head_dim + self.num_channels_latents = self.in_channels + + +@dataclass +class AnimaDiTConfig(DiTConfig): + arch_config: AnimaArchConfig = field(default_factory=AnimaArchConfig) + prefix: str = "anima" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/anima.py b/python/sglang/multimodal_gen/configs/pipeline_configs/anima.py new file mode 100644 index 000000000000..4226d2873860 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/anima.py @@ -0,0 +1,91 @@ +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass, field +from typing import Callable + +import torch + +from sglang.multimodal_gen.configs.models.dits.anima import AnimaDiTConfig +from sglang.multimodal_gen.configs.models.encoders.qwen3 import Qwen3TextConfig +from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ImagePipelineConfig, + ModelTaskType, +) + + +def anima_text_output(outputs, inputs): + return outputs.last_hidden_state * inputs.attention_mask.unsqueeze(-1) + + +@dataclass +class AnimaPipelineConfig(ImagePipelineConfig): + task_type: ModelTaskType = ModelTaskType.T2I + native_only_components: tuple[str, ...] = ("transformer", "text_conditioner") + dit_config: AnimaDiTConfig = field(default_factory=AnimaDiTConfig) + vae_config: QwenImageVAEConfig = field(default_factory=QwenImageVAEConfig) + text_encoder_configs: tuple = field(default_factory=lambda: (Qwen3TextConfig(),)) + text_encoder_precisions: tuple[str, ...] = ("bf16",) + preprocess_text_funcs: tuple[Callable | None, ...] = (None,) + postprocess_text_funcs: tuple[Callable, ...] = (anima_text_output,) + should_use_guidance: bool = False + enable_autocast: bool = False + vae_precision: str = "bf16" + + def tokenize_prompt(self, prompt, tokenizer, tok_kwargs): + if not 1 <= tok_kwargs.get("max_length", 512) <= 4096: + raise ValueError("Anima max_sequence_length must be between 1 and 4096") + inputs = tokenizer(prompt, **{**tok_kwargs, "padding": "longest"}) + if inputs.input_ids.shape[1] == 0: + inputs["input_ids"] = inputs.input_ids.new_zeros((len(prompt), 1)) + inputs["attention_mask"] = inputs.attention_mask.new_zeros((len(prompt), 1)) + return inputs + + def prepare_sigmas(self, sigmas, num_inference_steps): + return self._prepare_sigmas(sigmas, num_inference_steps) + + def get_latent_dtype(self, prompt_dtype): + return torch.float32 + + def expand_conditioning_to_sample_batch(self, batch): + count = batch.num_outputs_per_prompt + if count > 1: + batch.prompt_embeds = [ + x.repeat_interleave(count, dim=0) for x in batch.prompt_embeds + ] + if batch.do_classifier_free_guidance: + batch.negative_prompt_embeds = [ + x.repeat_interleave(count, dim=0) + for x in batch.negative_prompt_embeds + ] + return batch + + # the DiT shards patch tokens, leaving scheduler latents replicated + def shard_latents_for_sp(self, batch, latents): + return latents, False + + def get_pos_prompt_embeds(self, batch): + return batch.prompt_embeds[0] + + def get_neg_prompt_embeds(self, batch): + return batch.negative_prompt_embeds[0] + + def get_decode_scale_and_shift(self, device, dtype, vae): + config = self.vae_config.arch_config + std = torch.tensor(config.latents_std, device=device, dtype=dtype) + mean = torch.tensor(config.latents_mean, device=device, dtype=dtype) + return std.reciprocal().view(1, -1, 1, 1, 1), mean.view(1, -1, 1, 1, 1) + + def post_decoding(self, frames, server_args): + return frames.squeeze(2) + + +def register(): + from sglang.multimodal_gen.configs.sample.anima import AnimaSamplingParams + from sglang.multimodal_gen.registry import register_configs + + register_configs( + sampling_param_cls=AnimaSamplingParams, + pipeline_config_cls=AnimaPipelineConfig, + hf_model_paths=["circlestone-labs/Anima-Base-v1.0-Diffusers"], + model_detectors=[lambda name: name.lower() == "animamodularpipeline"], + ) diff --git a/python/sglang/multimodal_gen/configs/sample/anima.py b/python/sglang/multimodal_gen/configs/sample/anima.py new file mode 100644 index 000000000000..045fbec54a3c --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/anima.py @@ -0,0 +1,14 @@ +# SPDX-License-Identifier: Apache-2.0 +from dataclasses import dataclass + +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams + + +@dataclass +class AnimaSamplingParams(SamplingParams): + height: int = 1024 + width: int = 1024 + num_frames: int = 1 + num_inference_steps: int = 30 + guidance_scale: float = 4.0 + negative_prompt: str = "" diff --git a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py index d200e6f8685f..bd10bc51f1c1 100644 --- a/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py +++ b/python/sglang/multimodal_gen/runtime/cache/cache_dit_integration.py @@ -403,6 +403,10 @@ class CustomBlockAdapterSpec: # Custom BlockAdapter metadata for models absent from cache-dit's registry. _CUSTOM_BLOCK_ADAPTER_SPECS: dict[str, CustomBlockAdapterSpec] = { + "AnimaTransformer3DModel": CustomBlockAdapterSpec( + blocks_attr="transformer_blocks", + forward_pattern=ForwardPattern.Pattern_3, + ), "MingImageTransformer2DModel": CustomBlockAdapterSpec( blocks_attr="layers", forward_pattern=ForwardPattern.Pattern_3, diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py index 589411032574..78c6b790e1d9 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/adapter_loader.py @@ -1,3 +1,6 @@ +from sglang.multimodal_gen.configs.models.adapter.anima import ( + AnimaTextConditionerConfig, +) from sglang.multimodal_gen.configs.models.adapter.ltx_2_connector import ( LTX2ConnectorConfig, ) @@ -10,9 +13,10 @@ class AdapterLoader(PlainStateDictComponentLoader): - component_names = ["connectors", "duration_head"] + component_names = ["connectors", "duration_head", "text_conditioner"] config_classes = { + "text_conditioner": AnimaTextConditionerConfig, "connectors": LTX2ConnectorConfig, "duration_head": LTX2DurationHeadConfig, } diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py index 1a4f44107995..bf31f2b77140 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/component_loader.py @@ -939,7 +939,7 @@ def load_customized( class TokenizerLoader(ComponentLoader): """Loader for tokenizers.""" - component_names = ["tokenizer", "text_tokenizer"] + component_names = ["tokenizer", "text_tokenizer", "t5_tokenizer"] expected_library = "transformers" def load_customized( diff --git a/python/sglang/multimodal_gen/runtime/models/adapter/anima.py b/python/sglang/multimodal_gen/runtime/models/adapter/anima.py new file mode 100644 index 000000000000..a52566867e96 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/adapter/anima.py @@ -0,0 +1,129 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Anima's learned T5-token queries over Qwen3 text states.""" + +import torch +import torch.nn.functional as F +from torch import nn + +from sglang.multimodal_gen.configs.models.adapter.anima import ( + AnimaTextConditionerConfig, +) +from sglang.multimodal_gen.runtime.layers.attention import LocalAttention +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) + + +class AnimaConditionerAttention(nn.Module): + def __init__(self, config, context_dim, cross_attention=False): + super().__init__() + dim = config.model_dim + self.heads = config.num_attention_heads + self.head_dim = dim // self.heads + self.q_proj = nn.Linear(dim, dim, bias=False) + self.k_proj = nn.Linear(context_dim, dim, bias=False) + self.v_proj = nn.Linear(context_dim, dim, bias=False) + self.o_proj = nn.Linear(dim, dim, bias=False) + self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6) + self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6) + self.attn = LocalAttention( + self.heads, self.head_dim, is_cross_attention=cross_attention + ) + + def forward(self, x, context, mask, rope, context_rope): + q = self.q_norm(self.q_proj(x).unflatten(-1, (self.heads, self.head_dim))) + k = self.k_norm(self.k_proj(context).unflatten(-1, (self.heads, self.head_dim))) + v = self.v_proj(context).unflatten(-1, (self.heads, self.head_dim)) + cos, sin = rope + q1, q2 = q.chunk(2, -1) + q = q * cos + torch.cat((-q2, q1), -1) * sin + cos, sin = context_rope + k1, k2 = k.chunk(2, -1) + k = k * cos + torch.cat((-k2, k1), -1) * sin + return self.o_proj(self.attn(q, k, v, attn_mask=mask).flatten(2)) + + +class AnimaConditionerBlock(nn.Module): + def __init__(self, config): + super().__init__() + self.use_self_attention = config.use_self_attention + norm = nn.LayerNorm if config.use_layer_norm else nn.RMSNorm + eps = 1e-5 if config.use_layer_norm else 1e-6 + dim = config.model_dim + if self.use_self_attention: + self.norm_self_attn = norm(dim, eps=eps) + self.self_attn = AnimaConditionerAttention(config, dim) + self.norm_cross_attn = norm(dim, eps=eps) + self.cross_attn = AnimaConditionerAttention( + config, config.source_dim, cross_attention=True + ) + self.norm_mlp = norm(dim, eps=eps) + self.mlp = nn.Sequential( + nn.Linear(dim, int(dim * config.mlp_ratio)), + nn.GELU(), + nn.Linear(int(dim * config.mlp_ratio), dim), + ) + + def forward(self, x, source, target_mask, source_mask, rope, source_rope): + if self.use_self_attention: + normed = self.norm_self_attn(x) + x = x + self.self_attn(normed, normed, target_mask, rope, rope) + x = x + self.cross_attn( + self.norm_cross_attn(x), source, source_mask, rope, source_rope + ) + return x + self.mlp(self.norm_mlp(x)) + + +class AnimaTextConditioner(nn.Module, LayerwiseOffloadableModuleMixin): + layerwise_offload_dit_group_enabled = False + layer_names = ["blocks"] + + def __init__(self, config: AnimaTextConditionerConfig): + super().__init__() + self.config = config + self.embed = nn.Embedding(config.target_vocab_size, config.target_dim) + self.in_proj = ( + nn.Linear(config.target_dim, config.model_dim) + if config.target_dim != config.model_dim + else nn.Identity() + ) + self.blocks = nn.ModuleList( + [AnimaConditionerBlock(config) for _ in range(config.num_layers)] + ) + self.out_proj = nn.Linear(config.model_dim, config.target_dim) + self.norm = nn.RMSNorm(config.target_dim, eps=1e-6) + + def _rope(self, length, device, dtype): + dim = self.config.model_dim // self.config.num_attention_heads + inv = 1.0 / (10000.0 ** (torch.arange(0, dim, 2, device=device).float() / dim)) + freqs = torch.outer(torch.arange(length, device=device).float(), inv) + freqs = torch.cat((freqs, freqs), dim=-1)[None, :, None, :] + return freqs.cos().to(dtype), freqs.sin().to(dtype) + + def forward( + self, + source_hidden_states, + target_input_ids, + target_attention_mask, + source_attention_mask, + ): + x = self.in_proj(self.embed(target_input_ids).to(source_hidden_states.dtype)) + rope = self._rope(x.shape[1], x.device, x.dtype) + source_rope = self._rope(source_hidden_states.shape[1], x.device, x.dtype) + # additive -inf preserves SDPA's zero output for fully masked rows + target_mask = torch.zeros_like(target_attention_mask, dtype=x.dtype) + target_mask.masked_fill_(target_attention_mask == 0, float("-inf")) + source_mask = torch.zeros_like(source_attention_mask, dtype=x.dtype) + source_mask.masked_fill_(source_attention_mask == 0, float("-inf")) + for block in self.blocks: + x = block( + x, source_hidden_states, target_mask, source_mask, rope, source_rope + ) + x = self.norm(self.out_proj(x)) * target_attention_mask.unsqueeze(-1).to( + x.dtype + ) + return F.pad(x, (0, 0, 0, max(0, self.config.min_sequence_length - x.shape[1]))) + + +EntryClass = AnimaTextConditioner diff --git a/python/sglang/multimodal_gen/runtime/models/dits/anima.py b/python/sglang/multimodal_gen/runtime/models/dits/anima.py new file mode 100644 index 000000000000..e508040d78ea --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/anima.py @@ -0,0 +1,329 @@ +# Copyright 2025 The NVIDIA Team and The HuggingFace Team. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Native Cosmos Predict2 image transformer used by Anima.""" + +import math + +import torch +from diffusers.models.embeddings import Timesteps +from torch import nn + +from sglang.multimodal_gen.configs.models.dits.anima import AnimaDiTConfig +from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in +from sglang.multimodal_gen.runtime.distributed import get_tp_world_size +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + gather_seq, + shard_like, + shard_seq, + tail_attn_meta, +) +from sglang.multimodal_gen.runtime.layers.attention import USPAttention +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm +from sglang.multimodal_gen.runtime.layers.linear import ( + ColumnParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT + + +def _is_anima_block(name, module): + return is_module_list_entry_in(name, ("transformer_blocks",)) + + +class AnimaPatchEmbed(nn.Module): + def __init__(self, config): + super().__init__() + self.patch_size = config.patch_size + channels = config.in_channels + int(config.concat_padding_mask) + self.proj = ReplicatedLinear( + channels * math.prod(config.patch_size), config.hidden_size, bias=False + ) + + def forward(self, x): + b, c, t, h, w = x.shape + pt, ph, pw = self.patch_size + x = x.reshape(b, c, t // pt, pt, h // ph, ph, w // pw, pw) + x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4).flatten(1, 3) + return self.proj(x)[0] + + +class AnimaRotaryEmbedding(nn.Module): + def __init__(self, config): + super().__init__() + d = config.attention_head_dim + self.dims = (d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)) + self.scales = config.rope_scale + self.patch_size = config.patch_size + + def forward(self, x): + shape = [size // patch for size, patch in zip(x.shape[-3:], self.patch_size)] + freqs = [] + for axis, (dim, scale) in enumerate(zip(self.dims, self.scales)): + theta = 10000.0 * scale ** (dim / (dim - 2)) + inv = 1.0 / ( + theta ** (torch.arange(0, dim, 2, device=x.device).float() / dim) + ) + angles = torch.outer( + torch.arange(shape[axis], device=x.device).float(), inv + ) + view = [1, 1, 1, dim // 2] + view[axis] = shape[axis] + freqs.append(angles.view(view).expand(*shape, dim // 2)) + angles = torch.cat(freqs * 2, dim=-1).flatten(0, 2) + return angles.cos(), angles.sin() + + +class AnimaTimeEmbedding(nn.Module): + def __init__(self, hidden_size): + super().__init__() + self.time_proj = Timesteps( + hidden_size, flip_sin_to_cos=True, downscale_freq_shift=0 + ) + self.t_embedder = nn.Module() + self.t_embedder.linear_1 = ReplicatedLinear( + hidden_size, hidden_size, bias=False + ) + self.t_embedder.linear_2 = ReplicatedLinear( + hidden_size, 3 * hidden_size, bias=False + ) + self.norm = RMSNorm(hidden_size, eps=1e-6) + + def forward(self, timestep, dtype): + x = self.time_proj(timestep).to(dtype) + emb = self.t_embedder.linear_1(x)[0] + emb = self.t_embedder.linear_2(torch.nn.functional.silu(emb))[0] + return emb, self.norm(x) + + +class AnimaAdaLayerNorm(nn.Module): + def __init__(self, hidden_size, rank, gated=True): + super().__init__() + self.gated = gated + self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) + self.linear_1 = ReplicatedLinear(hidden_size, rank, bias=False) + self.linear_2 = ReplicatedLinear( + rank, (3 if gated else 2) * hidden_size, bias=False + ) + + def forward(self, x, embedded_timestep, temb): + modulation = self.linear_1(torch.nn.functional.silu(embedded_timestep))[0] + modulation = self.linear_2(modulation)[0] + modulation = modulation + temb[..., : modulation.shape[-1]] + values = modulation.unsqueeze(1).chunk(3 if self.gated else 2, dim=-1) + out = self.norm(x) * (1 + values[1]) + values[0] + return (out, values[2]) if self.gated else out + + +class AnimaAttention(nn.Module): + def __init__(self, config, cross_attention, quant_config, prefix): + super().__init__() + self.heads = config.num_attention_heads // get_tp_world_size() + self.head_dim = config.attention_head_dim + size = config.hidden_size + context_size = config.text_embed_dim if cross_attention else size + self.to_q = ColumnParallelLinear( + size, size, bias=False, quant_config=quant_config, prefix=f"{prefix}.to_q" + ) + self.to_k = ColumnParallelLinear( + context_size, + size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.to_k", + ) + self.to_v = ColumnParallelLinear( + context_size, + size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.to_v", + ) + self.to_out = nn.ModuleList( + [ + RowParallelLinear( + size, + size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.to_out.0", + ) + ] + ) + self.norm_q = RMSNorm(self.head_dim, eps=1e-6) + self.norm_k = RMSNorm(self.head_dim, eps=1e-6) + self.attn = USPAttention( + self.heads, + self.head_dim, + is_cross_attention=cross_attention, + skip_sequence_parallel=cross_attention, + prefix=prefix, + ) + + def forward(self, x, context=None, rope=None, attn_mask_meta=None): + context = x if context is None else context + q = self.to_q(x)[0].unflatten(-1, (self.heads, self.head_dim)) + k = self.to_k(context)[0].unflatten(-1, (self.heads, self.head_dim)) + v = self.to_v(context)[0].unflatten(-1, (self.heads, self.head_dim)) + q, k = self.norm_q(q), self.norm_k(k) + if rope is not None: + cos, sin = (r[None, :, None, :] for r in rope) + # Cosmos uses split-half RoPE in fp32, then rounds once + q1, q2 = q.chunk(2, dim=-1) + k1, k2 = k.chunk(2, dim=-1) + q = (q.float() * cos + torch.cat((-q2, q1), -1).float() * sin).to(q.dtype) + k = (k.float() * cos + torch.cat((-k2, k1), -1).float() * sin).to(k.dtype) + out = self.attn(q, k, v, attn_mask_meta=attn_mask_meta).flatten(2) + return self.to_out[0](out)[0] + + +class AnimaFeedForward(nn.Module): + def __init__(self, config, quant_config, prefix): + super().__init__() + size = config.hidden_size + inner = int(size * config.mlp_ratio) + projection = nn.Module() + projection.proj = ColumnParallelLinear( + size, + inner, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.net.0.proj", + ) + self.net = nn.ModuleList( + [ + projection, + nn.Identity(), + RowParallelLinear( + inner, + size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.net.2", + ), + ] + ) + + def forward(self, x): + x = torch.nn.functional.gelu(self.net[0].proj(x)[0]) + return self.net[2](x)[0] + + +class AnimaTransformerBlock(nn.Module): + def __init__(self, config, quant_config, prefix): + super().__init__() + size, rank = config.hidden_size, config.adaln_lora_dim + self.norm1 = AnimaAdaLayerNorm(size, rank) + self.attn1 = AnimaAttention(config, False, quant_config, f"{prefix}.attn1") + self.norm2 = AnimaAdaLayerNorm(size, rank) + self.attn2 = AnimaAttention(config, True, quant_config, f"{prefix}.attn2") + self.norm3 = AnimaAdaLayerNorm(size, rank) + self.ff = AnimaFeedForward(config, quant_config, f"{prefix}.ff") + + def forward( + self, + hidden_states, + encoder_hidden_states, + embedded_timestep, + temb, + rope, + attn_mask_meta=None, + ): + x, gate = self.norm1(hidden_states, embedded_timestep, temb) + hidden_states = hidden_states + gate * self.attn1( + x, rope=rope, attn_mask_meta=attn_mask_meta + ) + x, gate = self.norm2(hidden_states, embedded_timestep, temb) + hidden_states = hidden_states + gate * self.attn2( + x, context=encoder_hidden_states + ) + x, gate = self.norm3(hidden_states, embedded_timestep, temb) + return hidden_states + gate * self.ff(x) + + +class AnimaTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): + _aliases = ["CosmosTransformer3DModel"] + _fsdp_shard_conditions = [_is_anima_block] + _compile_conditions = [_is_anima_block] + layer_names = ["transformer_blocks"] + param_names_mapping = {} + reverse_param_names_mapping = {} + + def __init__(self, config: AnimaDiTConfig, hf_config, quant_config=None, **kwargs): + super().__init__(config=config, hf_config=hf_config) + arch = self.config + if ( + arch.extra_pos_embed_type + or arch.use_crossattn_projection + or arch.img_context_dim_in + ): + raise ValueError( + "Only Anima's Cosmos image transformer architecture is supported" + ) + if arch.num_attention_heads % get_tp_world_size(): + raise ValueError("Anima attention heads must be divisible by TP size") + self.hidden_size = arch.hidden_size + self.num_attention_heads = arch.num_attention_heads + self.num_channels_latents = arch.in_channels + self.patch_embed = AnimaPatchEmbed(arch) + self.rope = AnimaRotaryEmbedding(arch) + self.time_embed = AnimaTimeEmbedding(arch.hidden_size) + self.transformer_blocks = nn.ModuleList( + [ + AnimaTransformerBlock(arch, quant_config, f"transformer_blocks.{i}") + for i in range(arch.num_layers) + ] + ) + self.norm_out = AnimaAdaLayerNorm( + arch.hidden_size, arch.adaln_lora_dim, gated=False + ) + self.proj_out = ReplicatedLinear( + arch.hidden_size, math.prod(arch.patch_size) * arch.out_channels, bias=False + ) + self.__post_init__() + + def forward(self, hidden_states, encoder_hidden_states, timestep, **kwargs): + encoder_hidden_states = encoder_hidden_states.to(hidden_states.dtype) + b, _, t, h, w = hidden_states.shape + pt, ph, pw = self.config.patch_size + if t % pt or h % ph or w % pw: + raise ValueError( + "Anima latent dimensions must be divisible by the patch size" + ) + rope = self.rope(hidden_states) + if self.config.concat_padding_mask: + hidden_states = torch.cat( + [hidden_states, hidden_states.new_zeros(b, 1, t, h, w)], dim=1 + ) + hidden_states = self.patch_embed(hidden_states) + hidden_states, shard = shard_seq(hidden_states) + rope = tuple(shard_like(r, shard, dim=0) for r in rope) + attn_mask_meta = tail_attn_meta(shard, b, hidden_states.device) + temb, embedded_timestep = self.time_embed( + timestep.to(hidden_states.dtype) / 1000, hidden_states.dtype + ) + for block in self.transformer_blocks: + hidden_states = block( + hidden_states, + encoder_hidden_states, + embedded_timestep, + temb, + rope, + attn_mask_meta, + ) + hidden_states = self.norm_out(hidden_states, embedded_timestep, temb) + hidden_states = self.proj_out(hidden_states)[0] + hidden_states = gather_seq(hidden_states, shard.orig_len) + # checkpoint output ordering is (ph, pw, pt, c), unlike the input patches + hidden_states = hidden_states.reshape( + b, t // pt, h // ph, w // pw, ph, pw, pt, -1 + ) + return hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5).reshape( + b, self.config.out_channels, t, h, w + ) + + +EntryClass = AnimaTransformer3DModel diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py index 6e7edb8d175a..6b5e80ef56bc 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/qwen3.py @@ -215,6 +215,10 @@ def _masked_causal_attention( k_item = k[batch_index : batch_index + 1] v_item = v[batch_index : batch_index + 1] + if valid_len == 0: + outputs.append(torch.zeros_like(q_item)) + continue + real_output = self.attn( q_item[:, :valid_len], k_item[:, :valid_len], @@ -323,6 +327,8 @@ class Qwen3ForCausalLM(TextEncoder): - FSDP sharding for CPU offload """ + _aliases = ["Qwen3Model"] + def __init__(self, config: Qwen3TextConfig) -> None: super().__init__(config) diff --git a/python/sglang/multimodal_gen/runtime/pipelines/anima_pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines/anima_pipeline.py new file mode 100644 index 000000000000..82db283ad366 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/anima_pipeline.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import LoRAPipeline +from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import ( + InputValidationStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.anima import ( + AnimaTextConditioningStage, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages.text_encoding import ( + TextEncodingStage, +) + + +class AnimaPipeline(LoRAPipeline, ComposedPipelineBase): + pipeline_name = "AnimaModularPipeline" + _required_config_modules = [ + "text_encoder", + "tokenizer", + "t5_tokenizer", + "text_conditioner", + "vae", + "transformer", + "scheduler", + ] + + def create_pipeline_stages(self, server_args): + self.add_stage(InputValidationStage()) + self.add_stage( + TextEncodingStage( + text_encoders=[self.get_module("text_encoder")], + tokenizers=[self.get_module("tokenizer")], + ) + ) + self.add_stage( + AnimaTextConditioningStage( + self.get_module("text_conditioner"), self.get_module("t5_tokenizer") + ) + ) + self.add_standard_timestep_preparation_stage() + self.add_standard_latent_preparation_stage() + self.add_standard_denoising_stage() + self.add_standard_decoding_stage() + + +EntryClass = AnimaPipeline diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/anima.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/anima.py new file mode 100644 index 000000000000..9373f02c496a --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/anima.py @@ -0,0 +1,81 @@ +# SPDX-License-Identifier: Apache-2.0 +import torch + +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.pipelines_core.stages.condition_encoding import ( + ConditionEncodingStage, +) + + +class AnimaTextConditioningStage(ConditionEncodingStage): + def __init__(self, conditioner, tokenizer): + super().__init__() + self.conditioner = conditioner + self.tokenizer = tokenizer + + def component_uses(self, server_args, stage_name=None): + return [ + ComponentUse(self._component_stage_name(stage_name), "text_conditioner") + ] + + @torch.no_grad() + def forward(self, batch, server_args): + prompts = batch.prompt if isinstance(batch.prompt, list) else [batch.prompt] + max_length = batch.max_sequence_length or 512 + if not 1 <= max_length <= 4096: + raise ValueError("Anima max_sequence_length must be between 1 and 4096") + with self.use_declared_component( + component_name="text_conditioner", module=self.conditioner + ) as conditioner: + self.conditioner = conditioner + batch.prompt_embeds = [ + self._condition( + conditioner, + prompts, + batch.prompt_embeds[0], + batch.prompt_attention_mask[0], + max_length, + ) + ] + if batch.do_classifier_free_guidance: + negative = [batch.negative_prompt] * len(prompts) + batch.negative_prompt_embeds = [ + self._condition( + conditioner, + negative, + batch.negative_prompt_embeds[0], + batch.negative_attention_mask[0], + max_length, + ) + ] + # Cosmos attends to conditioner padding; additional BCG padding is not valid + batch.prompt_attention_mask = None + batch.negative_attention_mask = None + batch.prompt_embeds_mask = None + batch.negative_prompt_embeds_mask = None + batch.prompt_seq_lens = [ + [batch.prompt_embeds[0].shape[1]] * batch.prompt_embeds[0].shape[0] + ] + if batch.do_classifier_free_guidance: + batch.negative_prompt_seq_lens = [ + [batch.negative_prompt_embeds[0].shape[1]] + * batch.negative_prompt_embeds[0].shape[0] + ] + return batch + + def _condition(self, conditioner, prompts, embeds, source_mask, max_length): + tokens = self.tokenizer( + prompts, + padding="longest", + truncation=True, + max_length=max_length, + return_tensors="pt", + ).to(embeds.device) + embeds = embeds.to(dtype=next(conditioner.parameters()).dtype) + with set_forward_context(current_timestep=0, attn_metadata=None): + return conditioner( + embeds, tokens.input_ids, tokens.attention_mask, source_mask + ) diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index bbb49c869096..702a5c7d84e7 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -190,7 +190,9 @@ def choices(cls) -> list[str]: BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset( { + "anima-base-v1.0-diffusers", "black-forest-labs/flux.1-dev", + "circlestone-labs/anima-base-v1.0-diffusers", "comfy-org/ideogram-4", "efficient-large-model/sana1.5_1.6b_1024px_diffusers", "efficient-large-model/sana-video_2b_480p_diffusers", @@ -237,6 +239,7 @@ def choices(cls) -> list[str]: BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset( { + "AnimaPipelineConfig", "FluxPipelineConfig", "GlmImagePipelineConfig", "Ideogram4PipelineConfig", @@ -793,7 +796,7 @@ def _adjust_breakable_cuda_graph_support(self): return logger.warning( - "[Diffusion BCG] disabled for %s: only FLUX.1-dev, Ideogram-4, " + "[Diffusion BCG] disabled for %s: only Anima Base v1.0, FLUX.1-dev, Ideogram-4, " "jdopensource/JoyAI-Echo, Lightricks/LTX-2, LongCat-Image, " "MiniMax-H3, Qwen/Qwen-Image, Qwen/Qwen-Image-2512, " "Qwen/Qwen-Image-2.1, SANA1.5, " diff --git a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py index 60695e2639e0..7eee4260b1ea 100644 --- a/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/hf_diffusers_utils.py @@ -149,13 +149,12 @@ def _is_weight_bearing_diffusers_component(key: str, value: Any) -> bool: def _get_declared_weight_component_dirs(model_path: str) -> list[str]: - model_index_path = os.path.join(model_path, "model_index.json") + model_index_path = _local_model_index_path(model_path) if not os.path.exists(model_index_path): return [] try: - with open(model_index_path) as f: - model_index = json.load(f) + model_index = _read_model_index(model_index_path) except Exception as exc: logger.warning( "Failed to read model_index.json at %s: %s", model_index_path, exc @@ -374,13 +373,12 @@ def _ci_validate_diffusers_model(model_path: str) -> tuple[bool, bool]: def _verify_diffusers_model_complete(path: str) -> bool: """Check if a diffusers model directory has all required component subdirectories.""" - config_path = os.path.join(path, "model_index.json") + config_path = _local_model_index_path(path) if not os.path.exists(config_path): return False try: - with open(config_path) as config_file: - model_index = json.load(config_file) + model_index = _read_model_index(config_path) except Exception as exc: logger.warning("Failed to read model_index.json at %s: %s", config_path, exc) return False @@ -683,6 +681,29 @@ def maybe_download_lora( return target +def _local_model_index_path(model_path: str) -> str: + standard = os.path.join(model_path, "model_index.json") + if os.path.isfile(standard): + return standard + return os.path.join(model_path, "modular_model_index.json") + + +def _read_model_index(path: str) -> dict[str, Any]: + with open(path) as f: + config = json.load(f) + if os.path.basename(path) == "modular_model_index.json": + config.pop("_blocks_class_name", None) + for name, spec in config.items(): + if isinstance(spec, list) and len(spec) == 3 and isinstance(spec[2], dict): + # native loaders consume the component directories in this snapshot + if spec[2].get("subfolder", name) != name: + raise ValueError( + f"Modular component {name!r} must use its own subfolder" + ) + config[name] = spec[:2] + return config + + def verify_model_config_and_directory(model_path: str) -> dict[str, Any]: """ Verify that the model directory contains a valid diffusers configuration. @@ -695,16 +716,16 @@ def verify_model_config_and_directory(model_path: str) -> dict[str, Any]: """ # Check for model_index.json which is required for diffusers models - config_path = os.path.join(model_path, "model_index.json") + config_path = _local_model_index_path(model_path) if not os.path.exists(config_path): raise ValueError( - f"Model directory {model_path} does not contain model_index.json. " + f"Model directory {model_path} does not contain model_index.json " + "or modular_model_index.json. " "Only HuggingFace diffusers format is supported." ) # Load the config - with open(config_path) as f: - config = json.load(f) + config = _read_model_index(config_path) # Verify diffusers version exists if "_diffusers_version" not in config: @@ -745,26 +766,37 @@ def verify_model_config_and_directory(model_path: str) -> dict[str, Any]: return cast(dict[str, Any], config) -def _resolve_remote_repo_model_index_path(model_name_or_path: str) -> str: +def _resolve_remote_repo_model_index_path( + model_name_or_path: str, filename: str = "model_index.json" +) -> str: """Return a local path to a remote repo's ``model_index.json``""" try: # Cache-aware: no local_dir, so the selected Hub reuses its cache and # revalidates the remote file when online. - return hf_hub_download(repo_id=model_name_or_path, filename="model_index.json") + return hf_hub_download(repo_id=model_name_or_path, filename=filename) except EntryNotFoundError: - # Repo exists but has no model_index.json (single-model repo); let the - # caller fall through to the single-model path. + if filename == "model_index.json": + return _resolve_remote_repo_model_index_path( + model_name_or_path, "modular_model_index.json" + ) raise except Exception as online_err: cached_path = None if not envs.SGLANG_USE_MODELSCOPE.get(): from huggingface_hub import try_to_load_from_cache - cached = try_to_load_from_cache( - repo_id=model_name_or_path, filename="model_index.json" + filenames = ( + (filename, "modular_model_index.json") + if filename == "model_index.json" + else (filename,) ) - if isinstance(cached, str) and os.path.exists(cached): - cached_path = cached + for candidate in filenames: + cached = try_to_load_from_cache( + repo_id=model_name_or_path, filename=candidate + ) + if isinstance(cached, str) and os.path.exists(cached): + cached_path = cached + break if cached_path is not None: logger.warning( "Could not fetch model_index.json for '%s' from the Hugging Face " @@ -815,8 +847,7 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]: model_index_path = _resolve_remote_repo_model_index_path(model_name_or_path) # Load the model_index.json - with open(model_index_path) as f: - config: dict[str, Any] = json.load(f) + config = _read_model_index(model_index_path) # Verify it has the required fields if "_class_name" not in config: diff --git a/python/sglang/multimodal_gen/test/server/test_server_anima.py b/python/sglang/multimodal_gen/test/server/test_server_anima.py new file mode 100644 index 000000000000..25c32836b02e --- /dev/null +++ b/python/sglang/multimodal_gen/test/server/test_server_anima.py @@ -0,0 +1,75 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Opt-in HTTP validation: SGLANG_ANIMA_TEST_MODEL=.""" + +import base64 +import io +import os +import sys + +import numpy as np +import pytest +from openai import OpenAI +from PIL import Image + +from sglang.multimodal_gen.test.server.test_server_utils import ServerManager +from sglang.multimodal_gen.test.test_utils import get_dynamic_server_port + +pytestmark = pytest.mark.skipif( + not os.environ.get("SGLANG_ANIMA_TEST_MODEL"), + reason="set SGLANG_ANIMA_TEST_MODEL for full-checkpoint HTTP tests", +) + + +@pytest.fixture(scope="module") +def server(): + manager = ServerManager( + model=os.environ["SGLANG_ANIMA_TEST_MODEL"], + port=get_dynamic_server_port(), + extra_args="--num-gpus 1 --performance-mode speed --attention-backend fa", + ) + context = manager.start() + try: + yield context + finally: + context.cleanup() + + +@pytest.mark.parametrize("size,steps,outputs", [(512, 4, 2), (1024, 30, 1)]) +def test_repeated_generation(server, size, steps, outputs): + with OpenAI( + api_key="EMPTY", base_url=f"http://127.0.0.1:{server.port}/v1" + ) as client: + requests = [] + for _ in range(2): + result = client.images.generate( + model=server.model, + prompt="masterpiece, best quality, safe, watercolor landscape, a quiet seaside village at sunset", + size=f"{size}x{size}", + n=outputs, + response_format="b64_json", + extra_body={ + "seed": 42, + "num_inference_steps": steps, + "guidance_scale": 4.0, + "negative_prompt": "", + "generator_device": "cpu", + "output_format": "png", + }, + ) + assert len(result.data) == outputs + images = [] + for item in result.data: + with Image.open(io.BytesIO(base64.b64decode(item.b64_json))) as image: + assert image.size == (size, size) + pixels = np.asarray(image.convert("RGB")) + assert pixels.std() > 5, "degenerate image" + images.append(pixels) + requests.append(images) + for first, second in zip(*requests): + np.testing.assert_array_equal(first, second) + if outputs > 1: + assert not np.array_equal(requests[0][0], requests[0][1]) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/python/sglang/multimodal_gen/test/unit/test_anima.py b/python/sglang/multimodal_gen/test/unit/test_anima.py new file mode 100644 index 000000000000..14c295632c22 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_anima.py @@ -0,0 +1,228 @@ +# SPDX-License-Identifier: Apache-2.0 +import json +import sys +from contextlib import nullcontext +from types import SimpleNamespace + +import pytest +import torch +from diffusers.models.transformers.transformer_cosmos import CosmosRotaryPosEmbed +from transformers import BatchEncoding + +from sglang.cli.utils import get_is_diffusion_model +from sglang.multimodal_gen.configs.models.dits.anima import AnimaArchConfig +from sglang.multimodal_gen.configs.pipeline_configs.anima import AnimaPipelineConfig +from sglang.multimodal_gen.configs.sample.anima import AnimaSamplingParams +from sglang.multimodal_gen.registry import get_model_info +from sglang.multimodal_gen.runtime.breakable_cuda_graph.prompt_padding import ( + pad_masked_prompt_kwargs, +) +from sglang.multimodal_gen.runtime.models.dits.anima import AnimaRotaryEmbedding +from sglang.multimodal_gen.runtime.models.encoders.qwen3 import Qwen3Attention +from sglang.multimodal_gen.runtime.models.registry import ModelRegistry +from sglang.multimodal_gen.runtime.pipelines.anima_pipeline import AnimaPipeline +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.anima import ( + AnimaTextConditioningStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils import hf_diffusers_utils + + +def test_registry_recognizes_official_and_local_anima(monkeypatch): + monkeypatch.setattr( + "sglang.multimodal_gen.registry.maybe_download_model_index", + lambda _: {"_class_name": "AnimaModularPipeline"}, + ) + get_model_info.cache_clear() + for model in ("circlestone-labs/Anima-Base-v1.0-Diffusers", "/models/custom-anima"): + info = get_model_info(model) + assert info.pipeline_cls is AnimaPipeline + assert info.pipeline_config_cls is AnimaPipelineConfig + assert info.sampling_param_cls is AnimaSamplingParams + get_model_info.cache_clear() + assert ( + ModelRegistry.resolve_model_cls("Qwen3Model")[0].__name__ == "Qwen3ForCausalLM" + ) + assert ( + ModelRegistry.resolve_model_cls("AnimaTextConditioner")[0].__name__ + == "AnimaTextConditioner" + ) + + +def test_empty_prompt_keeps_a_masked_qwen_token(): + def tokenizer(prompts, **kwargs): + assert kwargs["padding"] == "longest" + return BatchEncoding( + { + "input_ids": torch.empty(len(prompts), 0, dtype=torch.long), + "attention_mask": torch.empty(len(prompts), 0, dtype=torch.long), + } + ) + + inputs = AnimaPipelineConfig().tokenize_prompt([""], tokenizer, {}) + assert inputs.input_ids.shape == (1, 1) + assert not inputs.attention_mask.any() + + +def test_breakable_cuda_graph_stays_enabled_for_anima(): + args = SimpleNamespace( + model_id="circlestone-labs/Anima-Base-v1.0-Diffusers", + model_path="/models/Anima-Base-v1.0-Diffusers", + enable_breakable_cuda_graph=True, + pipeline_config=AnimaPipelineConfig(), + warmup_resolutions=["512x512"], + ) + args._is_breakable_cuda_graph_supported_model = lambda: ( + ServerArgs._is_breakable_cuda_graph_supported_model(args) + ) + ServerArgs._adjust_breakable_cuda_graph_support(args) + assert args.enable_breakable_cuda_graph + + +@pytest.mark.parametrize("length", [512, 600]) +def test_conditioning_does_not_allow_extra_bcg_padding(length): + embeds = torch.randn(1, length, 8) + mask = torch.ones(1, length, dtype=torch.bool) + batch = SimpleNamespace( + prompt="landscape", + negative_prompt="", + max_sequence_length=1024, + prompt_embeds=[embeds], + negative_prompt_embeds=[embeds], + prompt_attention_mask=[mask], + negative_attention_mask=[mask], + prompt_embeds_mask=[mask], + negative_prompt_embeds_mask=[mask], + do_classifier_free_guidance=True, + ) + stage = SimpleNamespace( + conditioner=None, + use_declared_component=lambda **kwargs: nullcontext(None), + _condition=lambda *args: embeds, + ) + AnimaTextConditioningStage.forward(stage, batch, None) + for masks in (batch.prompt_embeds_mask, batch.negative_prompt_embeds_mask): + kwargs = {"encoder_hidden_states": embeds, "encoder_hidden_states_mask": masks} + assert pad_masked_prompt_kwargs(kwargs, (1024,)) is kwargs + assert batch.prompt_seq_lens == batch.negative_prompt_seq_lens == [[length]] + + +def test_qwen_all_masked_row_never_calls_attention_with_empty_kv(): + def attention(*args): + raise AssertionError("all-masked rows have no valid KV") + + module = SimpleNamespace(attn=attention) + q = torch.randn(1, 1, 2, 8) + out = Qwen3Attention._masked_causal_attention(module, q, q, q, (0,)) + torch.testing.assert_close(out, torch.zeros_like(q)) + + +@pytest.mark.parametrize("height,width", [(4, 6), (16, 16), (10, 14)]) +def test_rope_matches_cosmos(height, width): + arch = AnimaArchConfig() + x = torch.zeros(1, 16, 1, height, width) + actual = AnimaRotaryEmbedding(arch)(x) + reference = CosmosRotaryPosEmbed( + arch.attention_head_dim, + max_size=arch.max_size, + patch_size=arch.patch_size, + rope_scale=arch.rope_scale, + )(x) + for a, b in zip(actual, reference): + torch.testing.assert_close(a, b, atol=0, rtol=0) + + +def test_latents_and_conditioning_preserve_sample_batch(): + config = AnimaPipelineConfig() + config.vae_config.post_init() + batch = SimpleNamespace( + height=512, + width=768, + num_outputs_per_prompt=2, + prompt_embeds=[torch.randn(2, 512, 8)], + negative_prompt_embeds=[torch.randn(2, 512, 8)], + do_classifier_free_guidance=True, + ) + original = batch.prompt_embeds[0].clone() + assert config.prepare_latent_shape(batch, 4, 1) == (4, 16, 1, 64, 96) + assert config.get_latent_dtype(torch.bfloat16) == torch.float32 + config.expand_conditioning_to_sample_batch(batch) + torch.testing.assert_close(batch.prompt_embeds[0], original.repeat_interleave(2, 0)) + assert batch.negative_prompt_embeds[0].shape == (4, 512, 8) + latents = torch.randn(4, 16, 1, 64, 96) + value, sharded = config.shard_latents_for_sp(batch, latents) + assert value is latents and not sharded + + +def test_modular_index_uses_existing_component_loaders(tmp_path, monkeypatch): + index = { + "_class_name": "AnimaModularPipeline", + "_blocks_class_name": "AnimaAutoBlocks", + "_diffusers_version": "0.39.0.dev0", + "transformer": [ + "diffusers", + "CosmosTransformer3DModel", + {"subfolder": "transformer"}, + ], + "t5_tokenizer": ["transformers", "T5Tokenizer", {"subfolder": "t5_tokenizer"}], + } + path = tmp_path / "modular_model_index.json" + path.write_text(json.dumps(index)) + assert get_is_diffusion_model(str(tmp_path)) + (tmp_path / "transformer").mkdir() + (tmp_path / "transformer" / "diffusion_pytorch_model.safetensors").touch() + (tmp_path / "t5_tokenizer").mkdir() + config = hf_diffusers_utils.verify_model_config_and_directory(str(tmp_path)) + assert config["transformer"] == index["transformer"][:2] + assert "_blocks_class_name" not in config + assert hf_diffusers_utils._verify_diffusers_model_complete(str(tmp_path)) + calls = [] + + def download(repo_id, filename): + calls.append(filename) + if filename != "modular_model_index.json": + raise hf_diffusers_utils.EntryNotFoundError("not found") + return str(path) + + monkeypatch.setattr(hf_diffusers_utils, "hf_hub_download", download) + remote = hf_diffusers_utils.maybe_download_model_index("test/anima") + assert remote["transformer"] == config["transformer"] + assert calls[-2:] == ["model_index.json", "modular_model_index.json"] + (tmp_path / "model_index.json").write_text( + json.dumps( + { + "_class_name": "Legacy", + "_diffusers_version": "0.37.0", + "transformer": index["transformer"][:2], + } + ) + ) + assert ( + hf_diffusers_utils.verify_model_config_and_directory(str(tmp_path))[ + "_class_name" + ] + == "Legacy" + ) + + +def test_modular_index_works_from_offline_cache(tmp_path, monkeypatch): + path = tmp_path / "modular_model_index.json" + path.write_text("{}") + + def offline(**kwargs): + raise hf_diffusers_utils.RequestsConnectionError("offline") + + monkeypatch.setattr(hf_diffusers_utils, "hf_hub_download", offline) + monkeypatch.setattr( + "huggingface_hub.try_to_load_from_cache", + lambda repo_id, filename: ( + str(path) if filename == "modular_model_index.json" else None + ), + ) + assert hf_diffusers_utils._resolve_remote_repo_model_index_path( + "test/anima" + ) == str(path) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/python/sglang/multimodal_gen/test/unit/test_anima_cuda.py b/python/sglang/multimodal_gen/test/unit/test_anima_cuda.py new file mode 100644 index 000000000000..1ac22cfa735d --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_anima_cuda.py @@ -0,0 +1,122 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Small, checkpoint-free numerical contracts against the Cosmos reference.""" + +import sys + +import pytest +import torch +from diffusers.models.transformers.transformer_cosmos import CosmosTransformer3DModel + +from sglang.multimodal_gen.configs.models.adapter.anima import ( + AnimaTextConditionerArchConfig, + AnimaTextConditionerConfig, +) +from sglang.multimodal_gen.configs.models.dits.anima import AnimaDiTConfig +from sglang.multimodal_gen.configs.pipeline_configs.anima import AnimaPipelineConfig +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + cleanup_dist_env_and_memory, + maybe_init_distributed_environment_and_model_parallel, +) +from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context +from sglang.multimodal_gen.runtime.models.adapter.anima import AnimaTextConditioner +from sglang.multimodal_gen.runtime.models.dits.anima import AnimaTransformer3DModel +from sglang.multimodal_gen.runtime.server_args import ServerArgs, 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", autouse=True) +def cleanup_owned_process_group(): + initialized = torch.distributed.is_initialized() + yield + if not initialized: + cleanup_dist_env_and_memory() + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("batch_size", [1, 2]) +@torch.no_grad() +def test_transformer_matches_cosmos(dtype, batch_size, monkeypatch): + kwargs = dict( + in_channels=4, + out_channels=4, + num_attention_heads=4, + attention_head_dim=32, + num_layers=2, + mlp_ratio=2, + text_embed_dim=32, + adaln_lora_dim=16, + patch_size=(1, 2, 2), + max_size=(128, 240, 240), + rope_scale=(1.0, 4.0, 4.0), + concat_padding_mask=True, + extra_pos_embed_type=None, + use_crossattn_projection=False, + ) + config = AnimaDiTConfig() + config.update_model_arch(kwargs) + args = ServerArgs( + model_path="circlestone-labs/Anima-Base-v1.0-Diffusers", + num_gpus=1, + pipeline_config=AnimaPipelineConfig(dit_config=config), + attention_backend="torch_sdpa", + ) + monkeypatch.setattr(server_args, "_global_server_args", args) + ensure_distributed_env_defaults() + maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1) + torch.manual_seed(42) + reference = CosmosTransformer3DModel(**kwargs).cuda().to(dtype).eval() + model = AnimaTransformer3DModel(config, kwargs).cuda().to(dtype).eval() + model.load_state_dict(reference.state_dict(), strict=True) + x = torch.randn(batch_size, 4, 1, 8, 12, device="cuda", dtype=dtype) + context = torch.randn(batch_size, 7, 32, device="cuda", dtype=dtype) + t = torch.linspace(600, 900, batch_size, device="cuda") + padding = torch.zeros(1, 1, 8, 12, device="cuda", dtype=dtype) + expected = reference(x, t.to(dtype) / 1000, context, padding_mask=padding).sample + with set_forward_context(current_timestep=0, attn_metadata=None): + actual = model(x, context, t) + tolerance = 2e-5 if dtype == torch.float32 else 2e-2 + torch.testing.assert_close(actual, expected, atol=tolerance, rtol=tolerance) + + +@torch.no_grad() +def test_conditioner_masks_source_and_zero_pads_output(monkeypatch): + args = ServerArgs( + model_path="circlestone-labs/Anima-Base-v1.0-Diffusers", + num_gpus=1, + pipeline_config=AnimaPipelineConfig(), + attention_backend="torch_sdpa", + ) + monkeypatch.setattr(server_args, "_global_server_args", args) + config = AnimaTextConditionerConfig( + arch_config=AnimaTextConditionerArchConfig( + source_dim=32, + target_dim=32, + model_dim=32, + num_attention_heads=2, + num_layers=2, + target_vocab_size=16, + ) + ) + model = AnimaTextConditioner(config).cuda().bfloat16().eval() + source = torch.randn(2, 4, 32, device="cuda", dtype=torch.bfloat16) + ids = torch.tensor([[1, 2, 0], [1, 0, 0]], device="cuda") + target_mask = ids != 0 + source_mask = torch.tensor([[1, 1, 0, 0], [0, 0, 0, 0]], device="cuda") + changed = source.clone() + changed[source_mask == 0] += 10 + with set_forward_context(current_timestep=0, attn_metadata=None): + actual = model(source, ids, target_mask, source_mask) + expected = model(changed, ids, target_mask, source_mask) + torch.testing.assert_close(actual, expected, atol=0, rtol=0) + assert actual.shape == (2, 512, 32) + assert torch.isfinite(actual).all() + assert not actual[:, 3:].any() + assert not actual[:, :3][~target_mask].any() + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__])) diff --git a/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py b/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py index 2ed0184022b4..b364fdf66304 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py +++ b/python/sglang/multimodal_gen/test/unit/test_cache_dit_integration.py @@ -301,14 +301,22 @@ class TestBuildCustomBlockAdapter(unittest.TestCase): def test_builds_adapter_for_registered_class(self): module = _import_module_with_stub() blocks = ["block_0", "block_1"] - transformer = _make_transformer("ErnieImageTransformer2DModel", blocks) - - adapter = module._build_custom_block_adapter(transformer, has_separate_cfg=True) - - self.assertIsNotNone(adapter) - self.assertEqual(adapter.blocks, blocks) - self.assertEqual(adapter.forward_pattern, "Pattern_3") - self.assertTrue(adapter.has_separate_cfg) + for class_name, blocks_attr in ( + ("ErnieImageTransformer2DModel", "layers"), + ("MingImageTransformer2DModel", "layers"), + ("AnimaTransformer3DModel", "transformer_blocks"), + ): + with self.subTest(class_name=class_name): + transformer = type(class_name, (), {blocks_attr: blocks})() + adapter = module._build_custom_block_adapter( + transformer, has_separate_cfg=True + ) + + self.assertIsNotNone(adapter) + self.assertIs(adapter.transformer, transformer) + self.assertIs(adapter.blocks, blocks) + self.assertEqual(adapter.forward_pattern, "Pattern_3") + self.assertTrue(adapter.has_separate_cfg) def test_returns_none_for_unknown_class(self): module = _import_module_with_stub() diff --git a/scripts/ci/utils/diffusion/comparison_configs.json b/scripts/ci/utils/diffusion/comparison_configs.json index 7d5df4d83d37..a101dd7b43bc 100644 --- a/scripts/ci/utils/diffusion/comparison_configs.json +++ b/scripts/ci/utils/diffusion/comparison_configs.json @@ -3,6 +3,22 @@ "measurement_repeats": 3, "test_image_url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png", "cases": [ + { + "id": "anima_base_t2i_1024", + "model": "circlestone-labs/Anima-Base-v1.0-Diffusers", + "task": "text-to-image", + "prompt": "masterpiece, best quality, safe, watercolor landscape, a quiet seaside village at sunset", + "width": 1024, + "height": 1024, + "seed": 42, + "num_gpus": 1, + "frameworks": { + "sglang": { + "serve_args": "--warmup-mode server --performance-mode speed", + "extra_env": {} + } + } + }, { "id": "flux1_dev_t2i_1024", "model": "black-forest-labs/FLUX.1-dev",