Skip to content

compat: provide missing FA3 private forward - #103

Merged
yeahdongcn merged 2 commits into
mainfrom
xd/musa60103-fa3-private-forward-shim
Aug 17, 2026
Merged

yeahdongcn merged 2 commits into
mainfrom
xd/musa60103-fa3-private-forward-shim

Conversation

@yeahdongcn

@yeahdongcn yeahdongcn commented Aug 17, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

  • add flash_attn_interface._flash_attn_forward only when the provider omits it
  • adapt the public flash_attn_func(..., return_softmax_lse=True) result to the
    low-level (output, softmax_lse, S_dmask, rng_state) inference contract
  • preserve a provider's native private implementation unchanged and keep the
    patch idempotent

Motivation

vLLM-Omni Ring Attention imports _flash_attn_forward because it needs the
softmax LSE for partial-attention accumulation. The current MATE compatibility
module exposes the equivalent public output+LSE API but not that private
symbol, so the downstream FA3 availability probe evaluates false.

The adapter is deliberately conditional and version-independent. Once MATE or
another provider exposes _flash_attn_forward, torchada leaves the native
object untouched.

Safety boundaries

  • require the public callable to explicitly support return_softmax_lse
  • translate only the keyword subset used by the Ring FA3 inference path
  • fail closed if the public provider does not return at least output and LSE
  • return None for dropout-only auxiliary values unavailable from the public
    inference API

Validation

  • ruff check src/torchada/_patch.py tests/test_cuda_patching.py
  • focused compatibility tests: 2 passed, 212 deselected
  • full non-hardware test slice: 410 passed, 17 skipped, 41 deselected
  • wheel build: torchada-0.1.81-py3-none-any.whl
  • S5000 downstream smoke in
    vllm-omni:minimax-h3-20260815@sha256:23ae27867cd19ce848a27688dba262d0569c67ff3c3b599cc2a429f8ab184a8b:
    • installed this exact commit editable over the image's torchada 0.1.79
    • vLLM-MUSA platform-plugin import captured the shim before vLLM-Omni Ring
      globals; HAS_FA3=True
    • vLLM-Omni's unchanged FA3 selector resolved to
      ring_kernels.fa3_forward
    • BF16 [1,1024,14,128] output and LSE were finite; shim vs the public MATE
      output+LSE API had max absolute difference 0.0
    • real two-rank MCCL Ring Attention completed on both ranks with
      attn_type=fa3, ring_kernel=fa3_forward, and finite
      [1,512,14,128] output

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@yeahdongcn
yeahdongcn marked this pull request as ready for review August 17, 2026 08:36
Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@yeahdongcn
yeahdongcn merged commit ea7e436 into main Aug 17, 2026
yeahdongcn added a commit that referenced this pull request Oct 10, 2026
The English and Chinese READMEs have not kept up with the May-October work.
This documents what merged, moves the version-gated shims into one table, and
refreshes the measured numbers that had gone stale.

Feature table
- CUDA memory-pool APIs, `torch.cuda.streams`, CUDA-graph executable rotation,
  `torch.cuda._get_device_index`, `get_memory_info()`, and the FlashAttention
  provider shims (#61, #98, #103, #106, #108, #115)
- the "What Works" table goes back to one line per feature; the paragraph-sized
  `log_` / `isfinite` / `out_dtype` cells move into the new section below

New "torch_musa Compatibility" section
- one table of every version-gated shim with the release it is installed on:
  the four `< 2.11.0.post2` patches (#106, #113, #124), the `< 2.13.0`
  `mm`/`bmm` `out_dtype=` backport (#116), the stable-ABI header backport (#86,
  #96), asynchronous `isfinite` (#120), and `torch.cuda.streams` (#98)

New "Environment Variables" section
- the graph-rotation knobs (#72), `TORCHADA_PLATFORM`, the C++ operator-override
  switches (#61, #128), and the two variables that were already documented

Corrected and extended details
- torch.compile: FX `device` builtin (#124), Dynamo's device-index helper (#108),
  `MUSA_VISIBLE_DEVICES` mirroring (#106)
- C++ extensions: nested `<torch/cuda.h>` porting (#95), stable-ABI
  `STABLE_TORCH_LIBRARY_IMPL` rekeying and stream helpers (#100), torch 2.6+
  `include_paths`/`library_paths` signatures (#121), stale JIT build locks (#128)
- MoE tables are generated from checked-in recipes (#115)
- unsupported CUDA runtime APIs as no-ops (#65)
- Performance: replace the 0.1.94 / torch_musa 2.7.1 numbers with the checked-in
  0.1.95 / 2.11.0.post2 entry, and stop claiming every fast path is under 200ns
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant