Skip to content

Add SANA-WM camera-controlled image-to-video pipeline - #13881

Open
lawrence-cj wants to merge 63 commits into
huggingface:mainfrom
lawrence-cj:feat/sana-wm-diffusers-cleanup
Open

Add SANA-WM camera-controlled image-to-video pipeline#13881
lawrence-cj wants to merge 63 commits into
huggingface:mainfrom
lawrence-cj:feat/sana-wm-diffusers-cleanup

Conversation

@lawrence-cj

@lawrence-cj lawrence-cj commented Jun 7, 2026

Copy link
Copy Markdown
Contributor

What does this PR do?

Hi @sayakpaul @dg845 , Long time no see. Hoping your are doing great. ♥️

Adds SANA-WM, the camera-controlled image-to-video world model from NVIDIA + MIT HAN Lab, as a first-class diffusers pipeline and transformer. Given a first-frame image, a text prompt, and a camera trajectory (explicit c2w poses or a WASD/IJKL action-DSL string), the pipeline generates a video whose motion follows the requested camera path. Trained natively for minute-scale generation at 704×1280.

The pipeline runs in two stages:

  1. Stage 1 — SanaWMTransformer3DModel. A 1.6B-parameter bidirectional DiT with GDN-Triton linear attention and a UCPE camera-control branch; samples with an LTX-style flow-matching Euler scheduler at per-token timesteps. The first latent frame is the conditioning anchor.
  2. Stage 2 — SanaWMLTX2Refiner (optional). A chunk-causal AR refiner that wraps diffusers' LTX2VideoTransformer3DModel + LTX2TextConnectors + Gemma-3 text encoder. Processes 3 latent frames at a time with a sliding window of [source_sink + recent_history + active_block] K/V, so per-block compute is bounded and total refinement cost is linear in video length.

Both stages decode through AutoencoderKLLTX2Video.

Layout

src/diffusers/
├── models/transformers/
│   ├── transformer_sana_wm.py          # SanaWMTransformer3DModel + blocks + helpers
│   └── transformer_sana_wm_kernels.py  # fused Triton kernels + camera math
└── pipelines/sana_wm/
    ├── __init__.py
    ├── pipeline_sana_wm.py             # SanaWMPipeline
    ├── pipeline_output.py              # SanaWMPipelineOutput
    ├── refiner.py                      # SanaWMLTX2Refiner + RefinerChunkRunner
    └── cam_utils.py                    # action DSL, intrinsics, resize+crop, Plücker/raymap

scripts/sana_wm/convert_sana_wm_to_diffusers.py
docs/source/en/api/{pipelines/sana_wm.md, models/sana_wm_transformer3d.md}

Usage

import torch
from PIL import Image
from diffusers import SanaWMPipeline
from diffusers.utils import export_to_video

pipe = SanaWMPipeline.from_pretrained(
    "Efficient-Large-Model/SANA-WM_bidirectional-diffusers",
    torch_dtype=torch.bfloat16,
)
pipe.vae.to(torch.float32)
pipe.enable_model_cpu_offload()

out = pipe(
    image=Image.open("input.png").convert("RGB"),
    prompt="A car driving across a vast desert plain at golden hour.",
    action="w-80,jw-40,w-40",                    # WASD-style action DSL
    intrinsics=[800.0, 800.0, 845.0, 464.0],      # fx, fy, cx, cy in original-image pixels
    num_frames=161,
    num_inference_steps=60,
)
export_to_video(list(out.frames), "sana_wm.mp4", fps=16)

Demo

5-second sample (30 stage-1 steps + 3-step distilled AR refiner, official asset/sana_wm/demo_0 inputs, 704×1280 @ 16 fps) :

sana_wm_5s.mp4

Smoke tests

End-to-end on 1× H100 80GB with `enable_model_cpu_offload` and the official `asset/sana_wm/demo_0.{png,txt,_pose.npy,_intrinsics.npy}`:

Duration Frames Stage-1 (30 steps) Refiner (AR, 3 blocks) Output
5s 80 1:11 5:24 / step 525 KB
10s 160 1:11 28:55 (7 blocks) 1.4 MB
20s 320 1:57 ≈ 4 min / block (14) 3.2 MB
50s 800 5:33 30:46 (34 blocks) 6.3 MB

Checkpoint conversion

scripts/sana_wm/convert_sana_wm_to_diffusers.py --src Efficient-Large-Model/SANA-WM_bidirectional --dst /local/path converts the public release into a `from_pretrained`-loadable directory (VAE, Gemma-2 tokenizer + text_encoder, transformer, scheduler, refiner subfolders, top-level `model_index.json`).

Related

Paper: https://arxiv.org/abs/2605.15178

HaoyiZhu and others added 4 commits June 1, 2026 01:28
…line

Adds the public SANA-WM bidirectional camera-controlled image-to-video
model as a first-class diffusers pipeline + transformer. Layout mirrors
``sana_video``: the model lives under ``src/diffusers/models/transformers/``
as a near-single-file (kernels split off so the ``@triton.jit`` decorators
don't drown the model body); the pipeline lives under
``src/diffusers/pipelines/sana_wm/``.

Files added:

  src/diffusers/models/transformers/
  ├── transformer_sana_wm.py         # SanaWMTransformer3DModel + blocks + helpers
  └── transformer_sana_wm_kernels.py # fused Triton kernels + camera math

  src/diffusers/pipelines/sana_wm/
  ├── __init__.py
  ├── pipeline_sana_wm.py
  ├── pipeline_output.py
  ├── refiner.py
  └── cam_utils.py

Pipeline architecture:
* Stage 1: 1600M ``SanaWMTransformer3DModel`` DiT with bidirectional
  GDN-Triton linear attention + UCPE camera-control branch, LTX-style
  flow-matching Euler scheduler with per-token timesteps.
* Stage 2: LTX-2 sink-bidirectional Euler refiner (3 distilled sigma
  steps, reuses diffusers' ``LTX2VideoTransformer3DModel`` +
  ``LTX2TextConnectors`` + Gemma-3 text encoder).
* Decode through the LTX-2 VAE (``AutoencoderKLLTX2Video``).

One-line usage:

  pipe = SanaWMPipeline.from_pretrained(
      "Efficient-Large-Model/SANA-WM_bidirectional-diffusers",
      torch_dtype=torch.bfloat16,
  ).to("cuda")
  out = pipe(image=img, prompt="...", action="w-80,jw-40,w-40",
             intrinsics=[fx, fy, cx, cy])

End-to-end smoke test (stage-1 + refiner + VAE decode) passes on H100.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
…xport

transformer_sana_wm.py:
* License header switched to the "HuggingFace Team and SANA-WM Authors"
  style used by merged sana_video.
* Imports rewritten in stdlib -> third-party -> diffusers order; use
  diffusers `from ...utils import logging` instead of stdlib `logging`.
* Fix 9 `Optional[X]` annotations written as `X or None` (Python's `or`
  short-circuits and silently returns `X`).
* Fix two `assert (cond, msg)` tuple-asserts in PatchEmbedMS3D.forward
  that always pass (SyntaxWarning at import time).
* Remove duplicate `__all__` declarations (the second silently overwrote
  the first).
* Remove dead `reset_bn` (imports a nonexistent `packages.apps.utils`,
  would crash on call).
* Remove the duplicate `logger = logging.getLogger(__name__)` further
  down in the file.

transformer_sana_wm_kernels.py:
* License header normalized; collapse three duplicate triton/torch import
  blocks into one.

pipeline_sana_wm.py:
* License header normalized.
* `_decode_latents` now returns `(T, H, W, 3)` float in [0, 1], matching
  the diffusers convention used by `VideoProcessor`. Returning uint8
  silently broke `export_to_video`: it does `frame * 255` assuming float
  input, so uint8 overflows to `(-x) mod 256` and inverts colors.
* `__call__` converts to PIL/uint8 only when `output_type="pil"`.
* Intrinsics argument now accepts (4,), (F, 4), (3, 3), and (F, 3, 3)
  forms (auto-extracts fx, fy, cx, cy from a 3x3 K) and auto-trims to
  `num_frames` when a longer-than-needed trajectory is passed.
* Inline `retrieve_timesteps` with the standard `# Copied from
  diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps`
  marker, matching merged sana_video.
* Docstrings + EXAMPLE_DOC_STRING updated to reflect the new return type.

pipeline_output.py:
* Update `frames` field docstring to describe the new float [0, 1] return.

refiner.py, cam_utils.py, scripts/sana_wm/convert_sana_wm_to_diffusers.py:
* License headers normalized.

Docs:
* New `docs/source/en/api/pipelines/sana_wm.md` and
  `docs/source/en/api/models/sana_wm_transformer3d.md`, modeled on
  sana_video.md / sana_video_transformer3d.md, wired into
  `docs/source/en/_toctree.yml` under Models and Pipelines.

5s end-to-end smoke test (81 frames @ 16fps, 30 stage-1 steps + 3-step
LTX-2 refiner) passes on 1x H100 80GB with `enable_model_cpu_offload`.
Round-trip diff vs raw float frames is 2.06/255 mean (h264 lossy noise),
confirming the export_to_video fix.
…+ KV cache hooks)

The first cleanup pass only kept the legacy single-shot refiner path. That
path is what the model was *not* trained on — its docstring even says
"feeding the full sequence at once is out-of-distribution" — and its cost
is O(T^2) attention over the full latent volume, which made longer videos
unusable (~21 min per refiner step at 321 frames on an H100).

Port the chunk-causal AR mode from the upstream reference so the refiner
matches the training contract:

* `refine_latents` now defaults to `block_size=3, kv_max_frames=11`
  (the canonical AR recipe). Pass `block_size=None` to fall back to the
  legacy single-shot path.
* New `_refine_latents_ar` + `_RefinerChunkRunner` orchestrate the sliding
  window: pre-capture pre-RoPE sink K/V on `z_sana[:source_sink_frames]`
  at sigma=0, then for each `block_size`-frame chunk run a 3-step Euler
  with prefix `{sink_k_pre, sink_v, sink_pe, history_k, history_v}` and
  capture post-RoPE K/V to feed the next window. History is bounded to
  `kv_max_frames - source_sink_frames` so per-block compute is constant.
* New `_predict_x0_active_block` runs the transformer on the active block
  only (Q from active, K/V from prefix+active).
* New `_capture_block_kv` runs sigma=0 forward with a pre_rope/post_rope
  capture flag set on each `attn1`.
* New `_forward_video_only_with_rope` takes a pre-built RoPE so each block
  can use absolute frame positions in the source video.
* `_streaming_self_attention` extended with the `_kv_cache_capture`,
  `_tf_capture_kv`, `_tf_kv_prefix` hook contract that AR mode uses to
  inject and capture K/V on each block.
* New helpers: `_build_rotary_emb_for_absolute_positions`,
  `_set_kv_prefix_on_blocks`, `_clear_kv_prefix_on_blocks`,
  `_set_capture_flag_on_blocks`, `_collect_captured_kv_from_blocks`.
* `_encode_prompt` now also moves the Gemma-3 text encoder back to CPU
  after producing the embeds — otherwise it stays resident through the
  entire AR loop and gates how much GPU memory the refiner transformer
  has left.

Module-level docstring updated to document both modes; existing
single-shot path preserved verbatim.
…eemption)

The AR refiner is expensive (~3-5 min per block) and the refinement loop
ran end-to-end has no in-progress state to recover, so a SLURM preemption
mid-refinement loses all progress. With the canonical
``block_size=3, kv_max_frames=11`` setup, refining a 50s video is 34
blocks of work that has to make it through without preemption on a
backfill queue.

Add per-block atomic checkpointing:

* ``SanaWMLTX2Refiner.refine_latents(checkpoint_dir=Path)`` and
  ``_refine_latents_ar`` accept a directory. After each completed AR
  block, the AR loop writes ``checkpoint_dir/state.pt`` atomically
  (tmp + os.replace).
* The payload is ``{block_idx_done, n_blocks, sink_size, block_size,
  output_shape, output, runner_state}``. ``runner_state`` is a CPU snapshot
  of the runner's ``_sink_kv_pre``, ``_history_kv_post``,
  ``_history_frames`` and ``torch.Generator`` state.
* On entry, if ``state.pt`` exists with a compatible shape signature, the
  AR loop loads the persisted output tensor + runner state and resumes
  from ``block_idx_done + 1`` instead of recomputing from scratch.
* ``SanaWMPipeline.__call__(refiner_checkpoint_dir=...)`` plumbs the
  directory through to the refiner.

Checkpoint size: ~output_volume + sink_KV (~360MB for 50 layers) +
rolling history KV (~3-4GB at full capacity) — saved once per block,
total per-block save overhead ~10s on lustre.
@github-actions github-actions Bot added size/L PR with diff > 200 LOC documentation Improvements or additions to documentation models pipelines and removed size/L PR with diff > 200 LOC labels Jun 7, 2026
@github-actions github-actions Bot added the size/L PR with diff > 200 LOC label Jun 9, 2026
* CPU unit tests for cam_utils helpers (action DSL → c2w, intrinsics
  rescale-for-crop, resize+center-crop, snap_num_frames 8k+1 rounding).
* Public-surface registration tests (top-level diffusers symbols,
  SanaWMPipelineOutput dataclass shape, refiner signature has AR defaults
  + checkpoint_dir, pipeline __call__ accepts c2w/action/intrinsics/
  refiner_checkpoint_dir).
* @slow @require_torch_accelerator integration stub for an end-to-end I2V
  against the public checkpoint, currently @unittest.skip — wires up the
  nightly GPU path without exploding regular CI.

SanaWMTransformer3DModel has hardcoded depth/hidden_size/num_heads inside
its inner SanaMSVideoCamCtrl (not exposed through register_to_config), so
the usual PipelineTesterMixin small-config fast tests aren't applicable
without a transformer refactor (followup PR).
@github-actions github-actions Bot added the tests label Jun 9, 2026
@dg845
dg845 requested review from dg845 and yiyixuxu June 12, 2026 03:54
@dg845

dg845 commented Jun 16, 2026

Copy link
Copy Markdown
Collaborator

As a preliminary comment, would it be possible to use PyTorch ops instead of custom Triton kernels (or add pure PyTorch fallback paths) for now? We will work on supporting the custom kernels through kernels. CC @sayakpaul

@lawrence-cj

Copy link
Copy Markdown
Contributor Author

As a preliminary comment, would it be possible to use PyTorch ops instead of custom Triton kernels (or add pure PyTorch fallback paths) for now? We will work on supporting the custom kernels through kernels. CC @sayakpaul

Yes, love to do that.

…ttention

`transformer_sana_wm_kernels.py` previously did a hard `import triton`
at the top of the file. That blocked importing the SANA-WM transformer
on any environment without Triton (CPU-only, ROCm without Triton,
older Triton, etc.), even though the model has pure-PyTorch attention
classes for every `*Triton` variant.

Make Triton optional and have the dispatcher transparently fall back:

* Wrap `import triton` / `import triton.language as tl` in try/except.
  When unavailable, install a shim where `@triton.jit` is a no-op so
  the kernel function definitions still load (they just aren't compiled
  by Triton). Module-level `triton.X` / `tl.X` lookups return a
  self-shimming sentinel so signature parsing doesn't blow up either.
* Add `is_triton_available()` + `_require_triton(entry_point)`. The four
  Triton-backed entry points called by the model (`fused_qk_inv_rms`,
  `fused_bigdn_func`, `cam_prep_func`, `cam_scan_bidi_chunkwise`) now
  raise a clear RuntimeError on a Triton-less host with a hint to use
  the pure-PyTorch attention variants — but the dispatcher does this
  automatically (see below) so users shouldn't ever see it.
* Delete the leftover duplicate `import torch / triton / triton.language`
  block at line 262 (left over from the upstream port).
* Register `BidirectionalGDNUCPESinglePathLiteLA` in `ATTENTION_BLOCKS`
  so the fallback chain can find it.
* New `_resolve_attention_block(name, role)` walks the requested class's
  MRO at dispatch time. If Triton isn't usable AND the requested class
  name ends in `Triton`, route to the closest registered non-`Triton`
  ancestor (BidirectionalGDNUCPESinglePathLiteLABothTriton ->
  BidirectionalGDNUCPESinglePathLiteLA, etc.) and log a one-shot warning.
* Rewire both `SanaVideoMSCamCtrlBlock` dispatch sites to use
  `_resolve_attention_block` for the GDN+UCPE camera branch and the main
  attention branch (the `BidirectionalSoftmaxUCPESinglePathLiteLA` branch
  doesn't use Triton at all so it stays hard-coded).

Tests:
* `test_kernels_module_imports_with_triton_hidden` — reloads the kernels
  module with `sys.modules['triton'] = None` and verifies the module
  imports, `is_triton_available()` is False, and the pure-PyTorch helpers
  remain callable.
* `test_resolve_attention_block_cpu_fallback` — on a CPU-only host, the
  three `*Triton` attn types resolve to the correct non-Triton ancestor.
* `test_triton_entry_point_raises_clean_error_without_triton` — verifies
  the `_require_triton` guard yields a RuntimeError that mentions Triton.
@lawrence-cj

lawrence-cj commented Jun 16, 2026

Copy link
Copy Markdown
Contributor Author

Done in c0712d3f8 — Triton is now optional, with an automatic pure-PyTorch fallback at dispatch time. Mapping when Triton isn't usable:

Requested Falls back to
BidirectionalGDNTriton BidirectionalGDN
BidirectionalGDNUCPESinglePathLiteLATriton BidirectionalGDNUCPESinglePathLiteLA
BidirectionalGDNUCPESinglePathLiteLABothTriton BidirectionalGDNUCPESinglePathLiteLA

Triton remains the default on CUDA + Triton ≥ 3. CPU tests added under tests/pipelines/sana_wm/.

@lawrence-cj

Copy link
Copy Markdown
Contributor Author

@dg845 @yiyixuxu Gentle ping here.

lawrence-cj and others added 4 commits June 18, 2026 11:26
Three CI checks were failing on the PR:

1. `check_code_quality` (43 ruff errors): mix of unused imports / import
   sorting / E731 lambdas (auto-fixable) plus a handful of F821 dead-code
   references inherited from the upstream research codebase (`xformers.*`
   inside `if _xformers_available:` blocks, an undefined `BlockHook` type
   annotation, two `x_sa`/`mlp_out` references in a block forward whose
   live assignment was already overridden by subclasses). Ran `ruff check
   --fix --unsafe-fixes` + `ruff format`, fixed the type annotation
   manually, and added targeted `# noqa: F821` markers on the conditionally
   unreachable lines.

2. `check_torch_dependencies`: `transformer_sana_wm.py` hard-imported
   `einops`, `fla`, `timm`, `termcolor`. The minimum-deps CI environment
   doesn't have them, and diffusers' lazy loader rewrites `ModuleNotFoundError`
   as `RuntimeError` so `test_pipeline_imports` blew up. Wrapped each of
   the four optional imports in a try/except shim — `rearrange`/
   `ShortConvolution`/`DropPath`/`Attention_`/`Mlp` become placeholders
   that raise a clear `ImportError` on construction, `colored` falls back
   to plain text. Class bodies that subclass these still parse at module
   load, so `import diffusers.models.transformers.transformer_sana_wm`
   succeeds anywhere. Same treatment for the kernels file's
   `from einops import rearrange, repeat`.

3. `build_pr_documentation`: doc-builder imported `SanaWMTransformer3DModel`
   from `diffusers.models.transformers` (not the diffusers top level) and
   that subpackage's `__init__.py` was missing the entry. Added the import.
* `doc-builder style src/diffusers docs/source --max_len 119` rewraps
  docstrings in the six SANA-WM files (transformer, kernels, pipeline,
  refiner, output, cam_utils) to the repo-wide 119-column limit. No
  behaviour change — purely whitespace inside docstrings.
* `make fix-copies` regenerates `dummy_pt_objects.py` and
  `dummy_torch_and_transformers_objects.py` to add `DummyObject` stubs
  for the three new public classes (`SanaWMTransformer3DModel`,
  `SanaWMPipeline`, `SanaWMLTX2Refiner`), so `from diffusers import …`
  gives the standard "missing backend" message on installs without
  torch / transformers.

Verified: `make quality` passes (ruff check, ruff format check,
doc-builder style check_only, check_doc_toc). Test suite still
15 passed / 1 skipped.
@github-actions github-actions Bot added the utils label Jun 25, 2026
@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

@yiyixuxu

yiyixuxu commented Aug 25, 2026

Copy link
Copy Markdown
Collaborator

@lawrence-cj
for these 3 questions

Refiner composition — keep it nested, flatten Kandinsky-style, or split into two top-level pipelines?

have not reviewed in this PR

Attention — GDN linear attention doesn't fit the processor / dispatch_attention_fn pattern. Convert cross-attn + the softmax blocks only, and document GDN as an exception?

let's try to remove the timm dependency first in this round, don't worry about dispatch_attention_fn for now

Triton kernels — pure-PyTorch only in this PR (kernels as a follow-up), or publish to kernels-community and dispatch?

either one

Per @yiyixuxu's review. The three timm imports were all trivially replaceable:

* `Mlp` — both local subclasses fully overrode `forward` and only inherited
  `fc1` / `act` / `drop1` / `fc2` / `drop2`, so it is now a self-contained
  module and the redundant `class Mlp(Mlp)` wrapper is gone.
* `Attention` — `GDN` inherited only `qkv` and `proj` from it (it already
  defines its own `q_norm` / `k_norm`), so those two layers are declared
  inline and `GDN` subclasses `nn.Module`.
* `DropPath` — `drop_path` was hardcoded to `0.0`, so this was always
  `nn.Identity()`. Removed along with the rest of the plumbing.

No `dummy_timm_objects.py` is needed since the dependency is gone rather than
optional, and `setup.py` is untouched (other models still use timm).

State dict is unchanged (871/871 keys match the released checkpoint) and the
end-to-end GPU smoke gives byte-identical output (frame mean 0.5554).
@lawrence-cj

Copy link
Copy Markdown
Contributor Author

Thanks @yiyixuxutimm is gone in 64f24508c. All three imports turned out to be trivially replaceable:

  • Mlp — both local subclasses already overrode forward entirely and only inherited fc1/act/drop1/fc2/drop2, so it's now self-contained (and the redundant class Mlp(Mlp) wrapper is deleted).
  • AttentionGDN inherited only qkv and proj from it (it defines its own q_norm/k_norm), so those two layers are declared inline and it subclasses nn.Module.
  • DropPathdrop_path was hardcoded 0.0, so it was always nn.Identity(); removed with the rest of the plumbing.

No dummy_timm_objects.py needed since the dependency is removed rather than made optional, and setup.py is untouched (other models still use it). State dict is unchanged — 871/871 keys match the released checkpoint — and the end-to-end GPU smoke gives byte-identical output.

Leaving dispatch_attention_fn and the Triton question for later rounds as you suggested, and the refiner composition until you've had a chance to look.

One small thing whenever you get a moment: the workflow runs are still sitting on action_required for this fork, so nothing beyond the labeler has actually executed on the recent commits.

@yiyixuxu

Copy link
Copy Markdown
Collaborator

@lawrence-cj thanks, can you address all the inline review comments as well?

Per @yiyixuxu's inline review. Removes `transformer_sana_wm_kernels.py`
entirely (3234 lines) — the Triton kernels will be published to
`kernels-community` and wired back in a follow-up PR.

* Delete the three `*Triton` attention classes. Each only overrode `forward`
  to call a fused kernel; their pure-PyTorch parents (`BidirectionalGDN`,
  `BidirectionalGDNUCPESinglePathLiteLA`) compute the same result. The
  Triton-availability probe and the MRO-walking fallback go with them.
  The released `config.json` still names the `*Triton` variants, so
  `ATTENTION_BLOCKS` maps those strings onto the pure-PyTorch classes.
* Move the nine pure-PyTorch camera-math helpers that were reachable from
  the model into `transformer_sana_wm.py`; drop the rest with the kernels
  file (they were only reachable from the deleted Triton path).

Other review items:
* Use `get_activation()` from `..activations`; delete the local activation
  and normalization registries (`build_norm` was only ever called with
  `None`).
* Delete `GLUMBConv` and make `GLUMBConvTemp` self-contained. `ConvLayer` is
  slimmed and renamed `SanaWMConvLayer`, but has to stay a module — the
  checkpoint keys are `mlp.inverted_conv.conv.weight`, so collapsing it to a
  bare `nn.Conv2d` would drop the `.conv` level.
* Remove the closure-returning helpers the review flagged as hard to read
  and compile-hostile: `prepare_prope_fns` and friends now return plain
  tensors via `_prepare_ucpe_ray_transforms` + `_apply_ucpe_transform`, and
  the `_register_block` decorator factory is an explicit dict update.
* Remove all seven `@torch.compile` decorators — users can reach for
  `compile_repeated_blocks()` instead.
* Drop `partial` (pass `chunk_size` at the call site), `DWMlp`, the tuple
  helpers, `get_same_padding`, `modulate`, an unused `apply_rotary_emb`,
  and the validation of internal-only arguments.
* Tests: drop `SanaWMTritonFallbackTests` along with the mechanism it covered.

State dict is unchanged (871/871 keys) and the end-to-end GPU smoke passes
with correct video output.
…ly shim

Two more items from @yiyixuxu's inline review:

* Inline `t2i_modulate` at its four call sites.
* Delete `_IdentityForwardContiguousBackward` / `_contiguous_backward`. It is
  the identity in forward and only exists to hand a contiguous gradient to
  the backward pass, so it is dead weight in an inference-only port.

State dict unchanged (871/871 keys); GPU smoke output matches the previous
run exactly (frame mean 0.5560).
@lawrence-cj

lawrence-cj commented Aug 26, 2026

Copy link
Copy Markdown
Contributor Author

Thanks for the thorough pass @yiyixuxu — worked through all of them. transformer_sana_wm.py is 6555 → 4123 lines, and transformer_sana_wm_kernels.py (3234 lines) is gone.

Applied: Triton removed entirely (the *Triton config strings now map onto the pure-PyTorch classes, so existing checkpoints still load — I'll publish the kernels to kernels-community in a follow-up); timm dropped; get_activation(); GLUMBConv/DWMlp and the local registries deleted; all closure-returning helpers, @torch.compile decorators, partial, and the training-only inits removed.

Not applied — 4 hard blockers, not preferences:

Why
RMSNorm Shared class has no scale_factor; config sets y_norm_scale_factor=0.01. Also changes bf16 numerics. Can upstream scale_factor instead.
Timesteps/TimestepEmbedding Renames state-dict keys (t_embedder.mlp.0/2 vs linear_1/linear_2).
uniform-only chunking Config is first_chunk_plus_one and that path is live.
drop y_embedding Unused in forward, but it's one of the 871 checkpoint keys.

Each is 1 line to change if you'd rather take the checkpoint re-export — just say which.

Verified: state dict unchanged at 871/871 keys; GPU smoke on the public checkpoint gives correct video.

Workflow runs are still on action_required for this fork whenever you get a chance.

@dg845 dg845 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@lawrence-cj thanks for your patience! I have reviewed the refactored code. It would be helpful if you could run another self-review as described in #13881 (comment) after addressing the comments, as this will help speed up the review process.

Comment thread tests/pipelines/sana_wm/test_sana_wm.py
Comment thread tests/pipelines/sana_wm/test_sana_wm.py Outdated
Comment thread tests/pipelines/sana_wm/test_sana_wm.py Outdated
if mask is not None and mask.ndim == 2:
mask = (1 - mask.to(q.dtype)) * -10000.0
mask = mask[:, None, None].repeat(1, self.num_heads, 1, 1)
x = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Since MultiHeadCrossAttention uses the torch native F.scaled_dot_product_attention, I think we should follow the standard diffusers attention pattern and use dispatch_attention_fn here instead, as this will allow us to support different attention backends such as Flash Attention.

return hidden_states.permute(0, 2, 3, 1).reshape(batch_size, seq_len, channels)


class MultiHeadCrossAttention(nn.Module):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can you migrate MultiHeadCrossAttention to use the standard diffusers AttentionModuleMixin + attention processor design? For example, here is how Flux 2 implements the pattern:

class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
_default_processor_cls = Flux2AttnProcessor
_available_processors = [Flux2AttnProcessor, Flux2KVAttnProcessor]

class Flux2AttnProcessor:
_attention_backend = None
_parallel_config = None

Comment thread src/diffusers/pipelines/sana_wm/refiner.py Outdated
Comment on lines +451 to +460
for block in transformer.transformer_blocks:
hidden_states = _forward_video_block(
block=block,
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
video_rotary_emb=video_rotary_emb,
encoder_attention_mask=encoder_attention_mask,
n_context_tokens=n_context_tokens,
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If I understand correctly, _forward_video_block effectively implements a different forward method for LTX2VideoTransformerBlock. Given this, I think it would be better to implement the Sana-WM refiner transformer as a new model. We can use the # Copied from mechanism to sync the parts of the implementation that are shared with LTX2VideoTransformer3DModel, and add a new transformer block module whose forward follows the current _forward_video_block logic. CC @yiyixuxu

Comment thread src/diffusers/pipelines/sana_wm/cam_utils.py Outdated
Comment thread scripts/convert_sana_wm_to_diffusers.py
Comment thread scripts/sana_wm/convert_sana_wm_to_diffusers.py Outdated
Transformer:
* Stop stashing patch-grid shape on `self` during `forward` (`self.f/h/w`) and
  thread it through as locals; `unpatchify` now takes it explicitly.
* Replace the 7 `assert`s with `ValueError`s.
* Drop `attn_drop`/`proj_drop` from `MultiHeadCrossAttention` (training-only,
  and `attn_drop` was never applied) plus a stale training-era banner comment.

Pipeline / refiner:
* Don't mutate the components handed to the pipeline. The VAE tiling +
  framewise settings move to the docs, `padding_side="right"` is passed per
  tokenizer call, and the `.eval()` calls are gone (no `self.training`
  branches remain, and `from_pretrained` already returns eval-mode modules).
* Gate `cam_utils`' optional imports on `is_torchvision_available()` and a new
  `is_pi3_available()` helper.
* Annotate `SanaWMLTX2Refiner.__init__` and inline `_refine_latents_ar` into
  its single caller.

Conversion script:
* Move to `scripts/` alongside the other Sana converters.
* Raise on missing/unexpected keys instead of printing, so a bad mapping can't
  silently emit a broken transformer.

State dict unchanged (871/871 keys). GPU smoke on the public checkpoint is
byte-identical to the previous run (frame mean 0.5560), including with the VAE
settings applied by the caller rather than the pipeline.
The name implied this was Wan's rotary embedding, but it isn't: the per-axis
split is configurable through `fhw_dim`, and the frequencies stay complex in a
single `freqs` buffer instead of being split into real cos/sin buffers. So it
can't carry a `# Copied from`. Renamed, with a docstring recording why.

The buffer is `persistent=False`, so the state dict is unchanged (871/871).
… pytest

* `tests/models/transformers/test_models_transformer_sana_wm.py` — generated
  with `utils/generate_model_tests.py` and filled in, following
  `test_models_transformer_sana_video.py`. The tiny config sets
  `softmax_every_n=2` so one block exercises the GDN camera branch and the
  other the softmax variant. Dummy inputs supply the conditioning the forward
  requires: `encoder_attention_mask`, `(B, F, 20)` camera conditions, and
  `chunk_plucker`.
* `tests/pipelines/sana_wm/test_sana_wm.py` — rewritten in the pytest style of
  `tests/pipelines/sana_video/test_sana_video.py`: no `unittest`, bare
  asserts, module-level imports, and `parametrize` in place of loop-style
  cases (15 test functions become 31 cases).

`AttentionTesterMixin` is skipped because the model calls
`F.scaled_dot_product_attention` directly rather than going through a
diffusers attention processor.
@lawrence-cj

lawrence-cj commented Sep 2, 2026

Copy link
Copy Markdown
Contributor Author

Thanks @dg845 — pushed in 99d51ebd6, 82c902cf7, d44abcb7a. 15 of 21 threads resolved.

Highlights: forward no longer stashes the patch grid on self (a real hazard under concurrency/torch.compile, not just style); asserts → ValueError; the pipeline no longer mutates the components it's handed (VAE settings are documented instead); WanRotaryPosEmbedSanaWMRotaryPosEmbed, since it genuinely can't be # Copied from Wan; model tests added and pipeline tests migrated to pytest — 47 passed, 27 skipped.

Two flagged rather than done, both with a prerequisite:

  • Gradient checkpointing needs the block **kwargs bag (12 keys, 41 pass-through sites) flattened first, since diffusers' default checkpoint function is positional-only. Worth its own PR.
  • Manual offloading exists only because the refiner is nested; un-nesting it as you suggest makes the block disappear, so I'd rather do both together.

One correction: class SanaWMCamUtilsTests: would have silently disabled those tests — no pytest config, so python_classes is the default Test* and the unittest.TestCase base is the only reason they're collected today. Used TestSanaWMCamUtils.

Still yours + @yiyixuxu's call, since they reshape the public API: the dispatch_attention_fn migration, un-nesting the pipelines, and making the refiner transformer its own model. Happy to take all three.

Verified throughout: 871/871 state-dict keys, GPU smoke unchanged at frame mean 0.5560.

@dg845 dg845 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for the changes! I have left some follow up comments; can you also run make style and make quality to fix any code style errors?

@yiyixuxu could you take a look at the following comments?

Comment thread src/diffusers/models/transformers/transformer_sana_wm.py Outdated
Comment thread src/diffusers/models/transformers/transformer_sana_wm.py Outdated
Comment thread src/diffusers/pipelines/sana_wm/pipeline_sana_wm.py Outdated
Comment thread src/diffusers/pipelines/sana_wm/refiner.py Outdated
Comment on lines +156 to +159
# Free transformer GPU memory while we run the text encoder.
self.transformer.to("cpu")
empty_device_cache(device.type)
prompt_embeds, prompt_attention_mask = self._encode_prompt(prompt, device=device, dtype=dtype)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

As with SanaWMPipeline, we should use the standard model offloading utilities rather than manually implementing it here.

f"kv_prefix_per_layer has {len(kv_prefix_per_layer)} entries but transformer has {len(blocks)} blocks."
)
for block, prefix in zip(blocks, kv_prefix_per_layer):
block.attn1._tf_kv_prefix = prefix

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't think we should be modifying the LTX2VideoTransformer3DModel transformer here by setting attributes such as _tf_kv_prefix on its submodules. For KV caching, we can follow the Flux2KVCache pattern instead:

class Flux2KVCache:
"""Container for all layers' reference-token KV caches.
Holds separate cache lists for double-stream and single-stream transformer blocks.
"""

CC @yiyixuxu

Comment thread tests/models/transformers/test_models_transformer_sana_wm.py
lawrence-cj and others added 4 commits September 3, 2026 01:20
* `TimestepEmbedder.dtype` used `next(self.parameters()).dtype`, which reports
  the storage dtype under layerwise casting rather than the compute dtype. Use
  `get_parameter_dtype`, which is layerwise-casting aware. (The model-level
  `self.dtype` already routes through it via `ModelMixin`.)
* Stop reading submodule `.weight`/`.bias` inside `forward`. Group-offload
  hooks fire on a module's `forward`, so reading its parameters directly
  leaves them offloaded and trips a device mismatch. The frame-gate and
  output-gate helpers now call their submodules, and the fused camera QKV
  projection becomes three `q/k/v` calls -- algebraically the same as one
  GEMM over the concatenated weights.
* Don't mutate the tokenizer in `SanaWMLTX2Refiner._encode_prompt`; pass
  `padding_side="left"` per call, matching `SanaWMPipeline`.
* Add a `torch.compile` test with `recompile_limit=2` -- the repeated block
  compiles once per attention variant.

GPU smoke unchanged (frame mean 0.5560).
…instead of seed

Replace the `**kwargs` bag threaded through model -> block -> attention with
explicit keyword arguments. Only three runtime keys were ever read at the
leaves (`frame_valid_mask`, `precomputed_gates`, `ucpe_ray_transforms`); the
rest were forwarded and silently swallowed. `camera_embedding`,
`chunk_index`, `chunk_index_global` and `chunk_split_strategy` turned out to
be pure dead plumbing -- written into the per-block kwargs dicts and never
read by any attention or MLP forward -- so they are gone.

The pipeline also grows the standard diffusers arguments:

* `prompt_embeds` / `prompt_attention_mask` / `negative_prompt_embeds` /
  `negative_prompt_attention_mask` on `encode_prompt` and `__call__`.
* `seed` / `refiner_seed` are replaced by `generator` / `refiner_generator`.
  `generator=torch.Generator(device).manual_seed(42)` reproduces exactly what
  `seed=42` used to build, so results are unchanged.

State dict unchanged (871/871). CPU old-vs-new equality on a tiny config over
both attention variants is exact (`torch.equal`, max diff 0.0), and the GPU
smoke on the public checkpoint still gives frame mean 0.5560.
@yiyixuxu was right that the model only ever runs uniform chunking; my earlier
reply defending the strategies was wrong. The two sites that actually chunk
both call `normalize_chunk_index(None, T, chunk_size)` with three positional
arguments, so `chunk_split_strategy` always took its `"uniform"` default. The
configured `first_chunk_plus_one` reached the per-block kwargs dict and was
then swallowed by the attention forwards' `**kwargs` without ever being read.
Flattening that bag into explicit arguments is what surfaced it.

Also note the `chunk_size` those call sites use is `chunk_gdn_chunk_size`, a
different attribute from the `chunk_size` that was being threaded through.

So `chunk_index_from_chunk_size`, `normalize_chunk_index`,
`is_uniform_chunking` and `compute_chunk_sizes` are gone, the uniform
boundaries are inlined at both call sites, and `chunk_split_strategy` is
dropped from the model and block constructors.

The released `config.json` still carries the key; loading is unaffected
(`extract_init_dict` ignores it) but it logs an "not expected and will be
ignored" warning, so it should come out of the checkpoint config on the next
export.

State dict unchanged (871/871); GPU smoke still frame mean 0.5560, which
confirms those branches were never taken.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation models pipelines size/L PR with diff > 200 LOC tests utils

Projects

Status: In Progress

Development

Successfully merging this pull request may close these issues.

6 participants