Repository navigation
Conversation
…d_dim=256
- bump flash-attention submodule to jiayus-nvidia fork (435205c) which
supports head_dim=256 on SM100
- import `flash_attn.cute.*` (new fork's package layout) instead of the
monkey-patched `flash_attn_cute`
- fa4_utils: pick up `compile_cache` from the new top-level
`_bwd_preprocess` / `_bwd_postprocess_convert` functions (was an attr
on `_flash_attn_bwd` in the old fork)
- fa4: pass linear CSR masks via the dedicated
`linear_{q,k}_block_sparse_tensors` kwargs (the new fwd/bwd signature
separates them from `block_sparse_tensors`)
- calc_meta: keep `mask_block_cnt` at 3D `[B, H, num_blocks]` for the new
cute layout; split out `mask_definitions.flex_arbitrary_mask` import as
optional (not exported by the new fork)
- add focused 4-GPU CP test for head_dim=256 / bf16 / fa4 backend
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
The new FA cute backend's sparse-mask validator requires sparse_block_size_q == q_stage * tile_m. Previously magi always built the fwd mask with Q_BLOCK_SIZE = 2 * tile_m on sm100 and let FA infer the block size, which silently matched for MHA (seqlen_q <= 256 collapses to 1 q-block so FA inferred 128) but failed GQA with: "block size 128, which must be a multiple of 256" Plumb qhead_per_kvhead through CalcMeta -> FA4AttnArg, mirror FA's `q_stage = 2 if seqlen_q_packgqa > tile_m else 1` heuristic when building the sparse mask, and declare `block_size` explicitly on the LinearBlockSparseTensorsTorch so FA does not need to infer. Add a GQA hd256 test (8h/4kvh) alongside the existing MHA case. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
The triton-based output/lse correction kernel keeps both the input and output tiles in shared memory as fp32. At M_BLOCK=128 and head_dim=256 this needs ~256 KiB, which overflows even Blackwell's per-SM SMEM (232 KiB), raising triton.runtime.errors.OutOfResources: out of resource: shared memory, Required: 262144, Hardware limit: 232448 during the FA4 overlap-correction path. Halve M_BLOCK when the head_dim tile is at least 256. Extend tests/test_pipeline.py head_dim parametrize to include 256 so the full FA4 sweep on >=Hopper covers the new hd256 path (TestPipelineWithWorldSize4 with backend=fa4 now passes 96 cases in ~18 min on 4xB300). Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
The CalcMeta constructed by DistAttnSolver / DynamicAttnSolver was not populating qhead_per_kvhead, so users coming in through magi_attn_*_key (test_interface, test_pipeline, real Megatron-LM training) got the default 1, breaking the GQA + hd256 + FA4 path with: "block size 128, which must be a multiple of 256" Wire org_num_heads_q / org_num_heads_kv into CalcMeta from both solvers so the FA4 sparse-mask construction can mirror FA cute's q_stage heuristic correctly for any (qhq, qhkv) combo. Verified via TestPipelineWithWorldSize4.test_pipeline restricted to FA4 backend: 295 hd256 cases (MHA + GQA(8,2), fp16/bf16, varlen_full / varlen_block_causal / cp_mesh / etc.) all pass on 4xB300 in 16 min. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Apply project's pinned black 23.3.0 formatter to the new hd256 test (multi-line argument wrapping for the assert_close calls). Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
…only) The jiayus-nvidia fork dropped the top-level install.sh and the magi_to_hstu source dir. Switch to direct `pip install -e flash_attn/cute` for the JIT cute backend (which auto-supports head_dim=256 on sm100), and explicitly install create_block_mask_cuda from its new csrc/utils/ location. Trim the sm80/sm90 cutlass C++ build path entirely (not needed for our Blackwell target), drop the ARCH_ARG positional, and add a warning when magi_to_hstu_cuda is missing (its source no longer ships with the FA fork and must be supplied by the docker base image). Verified by running the script in the megatron-dev container then `tests/test_attn/test_dist_attn_hd256.py` (Ran 2 tests OK). Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
create_block_mask/setup.py auto-detects target arch via torch.cuda.get_device_capability(), which fails during docker image builds since no GPU is exposed. Detect that case in the install script and sed-patch the gencode flag to use MAGI_ATTENTION_CUDA_ARCH (default 100, i.e. sm_100) before pip install. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
End-to-end support for d_qk != d_v (e.g. DeepSeek-style d_qk=192, d_v=128)
in the distributed attention pipeline. Validated against the d192/k128
arbitrary-mask path on FA4 SM100 across all world sizes (ws=1..8) and
three sdpa attn_configs:
- sdpa_uneven_varlen_900
- sdpa_varlen_full_attn_1050
- sdpa_varlen_block_causal_960
Submodule
---------
flash-attention bumped from 435205c → 7045364 (jiayus-nvidia/flash-attention
magi_backend branch), which adds `Add SM100 d192v128 backward linear CSR
support`. The new SM100 backward kernel implements 2CTA + linear CSR for
hdim=192/v=128 and expects the K2Q block size to be (256, 256).
Public API
----------
`magi_attn_varlen_key` / `magi_attn_flex_key` (and the underlying
`init_dist_attn_runtime_*` helpers) gain an optional `head_dim_v` parameter:
magi_attn_flex_key(..., head_dim=192, head_dim_v=128)
Default `head_dim_v=None` is equivalent to head_dim_v=head_dim (symmetric
K/V), preserving backward compatibility. The key entry canonicalises
head_dim_v=head_dim → None so cache lookups don't split symmetric configs
across multiple LRU slots.
CommMeta (one place to decide packed_times)
-------------------------------------------
`CommMeta.head_dim_v` + `is_asymmetric_kv` property surface the user's
choice. `_init_a2av_based_grpcoll_args` reads `is_asymmetric_kv` once to
pick `packed_times` for the KV group_cast/reduce args:
- symmetric K/V: packed_times=2 (K + V along seqlen — the legacy path)
- asymmetric K/V: packed_times=1 (K, V will be fused along head_dim
inside _fetch_remote_kv, so only one tensor's worth of tokens flows)
`num_remote_kv_tokens_per_stage[stage]` is multiplied by the matching
factor at the same point. There is no longer any runtime-side lazy build
of "the other packed_times" view.
calc_meta
---------
On SM100, `bwd_sparse_tile_n` is doubled along with `bwd_sparse_tile_m`
to satisfy the new kernel's `BLOCK_SIZE=(256, 256)` requirement for the
K2Q linear CSR.
DistAttnRuntime (read once, hot path only consumes)
---------------------------------------------------
`DistAttnRuntime` reads `comm_meta.is_asymmetric_kv` in __init__ and
folds it into `concat_kv` / `concat_dkv` directly. No more flag flipping
at `DistAttnFunc.forward()` entry.
`_fetch_remote_kv` strategy for asymmetric K/V: concat K and V along the
head_dim (last) axis to form `[seqlen, h, d_k+d_v]`, then flow through
the existing fused-buffer logic. Token-partitioning semantics of the
comm primitive are preserved (it still partitions dim 0). On wait, the
fused buffer is split back into a (remote_k, remote_v) tuple by
hijacking `WorkWithPostProcessFn.post_process_fn` — no nested class, no
extra work wrapper.
`_reduce_partial_dkv` mirrors the strategy. A `_asym_fused_local_dkv`
cache (reset in `_reset_work_list`) shares one fused local-dkv buffer
across all overlap stages so multi-stage reduces accumulate into the
SAME tensor, mirroring how the symmetric concat_dkv=True path reuses
one partial_local_dkv across stages.
Bandwidth is exactly `seqlen * num_heads_kv * (d_k + d_v)` — same as
two independent K and V group_casts — but only one comm launch and one
set of comm args.
No new CPU-GPU sync points: all control flow is driven by Python ints /
bools / `is None` checks; tensor ops (cat / contiguous) are async kernel
launches on the existing CUDA stream.
Tests
-----
tests/test_pipeline.py: extend `FA4_D192_K128_ARBITRARY_CONFIGS` to cover
three sdpa configs above, and pass head_dim_v through to
`init_dist_attn_runtime_mgr` when head_dim_v != head_dim.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
The submodule pin (7045364, "Add magi_to_hstu CUDA utility") is on jiayus-nvidia/flash-attention's `magi_backend` branch; the old `arbitrary_mask_main_port` branch is no longer where new work lands. This keeps `git submodule update --remote` consistent with the pin direction. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
…module
magi_to_hstu now ships with the flash-attention submodule again as of
commit 7045364 ("Add magi_to_hstu CUDA utility") on the magi_backend
branch, so the install script should build it instead of just warning
that it's missing.
setup.py hard-codes -gencode for sm_80/90/100, so it builds correctly
during docker build without a GPU exposed — no need for the
torch.cuda.get_device_capability() fallback that create_block_mask uses.
The sanity-check warning is kept but reworded: a missing import after a
successful install means a silent build failure, not that the source is
elsewhere.
Also adds csrc/utils/magi_to_hstu to the MAGI_WHEEL_DIR sub-package list
so SCM-style wheel collection picks it up.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
demonatic
force-pushed
the
jager/ffa4_head_dim_256
branch
from
June 15, 2026 16:29
9954d53 to
903aa8e
Compare
Adds back the v1.1.1-style optional sm80/sm90 install: when the script is called with an arch list containing "sm80" or "sm90", it descends into flash-attention/hopper and runs `make install ARBITRARY=1 NUM_FUNC=1,3,5,7 HDIM128=1 SM8X=... SM90=...` to build the cutlass C++ kernels for Ampere and/or Hopper. When SM90 is requested without SM80, SM8X=0 is passed so the Ampere kernels are skipped (mirrors v1.1.1). Default invocation `bash scripts/install_flash_attn_cute.sh` is unchanged — sm100-only — keeping docker build time low for Blackwell images. bash scripts/install_flash_attn_cute.sh # sm100-only (default) bash scripts/install_flash_attn_cute.sh "sm80,sm90,sm100" # also build cutlass FA for Ampere/Hopper The sm100 cute path (flash_attn/cute + create_block_mask_cuda + magi_to_hstu_cuda) is unchanged and always installed. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
Collaborator
|
Since #331 already supports this, this PR can be closed. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
直接调用jiayu的支持head dim 256的ffa-fa4版本,目前test_pipeline.py就这2个head dim 256的case差一点精度:
