Skip to content

[Refactor] Chunk MAGI-2 Torch reference attention by query rows - #7282

Closed
yeahdongcn wants to merge 1 commit into
vllm-project:mainfrom
yeahdongcn:xd/magi2-chunked-attention
Closed

yeahdongcn wants to merge 1 commit into
vllm-project:mainfrom
yeahdongcn:xd/magi2-chunked-attention

Conversation

@yeahdongcn

Copy link
Copy Markdown
Contributor

Summary

Chunk the Torch reference path of MAGI-2 attention along the query axis, keeping every chunk's full key range.

  • Scores and probabilities cover at most query_chunk_size query rows (default 512) at a time during inference.
  • Preserve FP32 Q/K/V computation, GQA, softcap, sink logits, output dtype and sequence order.
  • Keep empty-query autograd dependencies as well as ordinary gradients. Invalid chunk sizes are rejected.

The bound concerns score/probability buffers only, not total memory: K/V storage and intermediates retained for autograd are not bounded by the query chunk size.

Scope

Two files only: the reference implementation and its tests. No backend selection changes, MUSA-specific path, new environment variables, native kernels, caches, RoPE, EP or sampler changes. Existing callers need no changes.

This is an independent extraction of the chunked-reference idea in #7156. It does not depend on the standalone FA3 adapter or on the MUSA consumer. The latter's pending numerical checks are not bypassed or reclassified here.

Validation

Commit: f644b504d7f1f2018fe8be21900845c876bd0d2d, one signed-off commit based on upstream 5378b77be.

  • Targeted pre-commit passed, including Ruff, mypy and test markers.
  • 44 CPU tests passed, 3 MUSA tests deselected: independent dense oracle, ragged/empty sequences, MHA/GQA, zero/one/multiple sinks, softcap, noncontiguous inputs, three floating dtypes, regular/empty-query gradients and existing native model tests.
  • A structural test observes the score operands: a 1025-query call with the default setting uses query row counts [512,512,1]. Separate ragged tests verify every chunk keeps the full per-sequence key range.
  • 3 MUSA tests passed on one 48-SM MTT S5000 (driver 5.2.0-server). FP32/BF16/FP16 chunk sizes 1 and 3 match the unchunked reference on the same GPU, with/without softcap and with GQA/two sinks. Tolerance remains rtol=1e-5, atol=2e-6.

Runtime: torch/torch_musa 2.11.0.post1+musa5.2.0, torchada 0.1.83, vLLM 0.28.0, vllm-musa 0.1.28. Exact source installed editable with --no-deps --no-build-isolation; MUSA entrypoint imports torchada first. No dependencies changed.

python -m pytest -q -o addopts= -m cpu \
  tests/diffusion/models/magi2/test_chunked_attention.py \
  tests/diffusion/models/magi2/test_native_preview.py

For MUSA, run the new test file with -m musa from a torchada-first entrypoint with one visible GPU.

Not run: CUDA hardware, GPU autograd, compile/graph, full MAGI-2 checkpoint/video, peak-memory or performance benchmarks. Structural buffer bounds and functional parity are not an E2E throughput/latency or full-resolution memory claim. This PR remains Draft.

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@hsliuustc0106 hsliuustc0106 added refactor refactoring for better code scalability and quality diffusion codes related to diffusion models labels Sep 8, 2026
@hsliuustc0106

Copy link
Copy Markdown
Collaborator

This PR touches tests/diffusion/, vllm_omni/diffusion/ (2 files). Based on CODEOWNERS coverage of the changed files, the most-related reviewers appear to be:

@fhfuih @Bounty-hunter @wtomin

Could one of you take a look when you get a chance? Thanks!

@yeahdongcn

Copy link
Copy Markdown
Contributor Author

Superseded by #8498, which carries this change as its second commit.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion codes related to diffusion models refactor refactoring for better code scalability and quality

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants