Skip to content

feat(ffn): add deterministic distributed Triton FFN for ROCm - #325

Open
frank-2077 wants to merge 78 commits into
RL-Align:testfrom
frank-2077:feat/rocm-strict-ffn
Open

feat(ffn): add deterministic distributed Triton FFN for ROCm#325
frank-2077 wants to merge 78 commits into
RL-Align:testfrom
frank-2077:feat/rocm-strict-ffn

Conversation

@frank-2077

@frank-2077 frank-2077 commented Aug 20, 2026

Copy link
Copy Markdown
Collaborator

Summary

This PR adds a ROCm-native deterministic distributed Qwen3 FFN implemented in
Triton. It supports FFN forward/backward across tensor parallelism (TP), context
parallelism (CP), and sequence parallelism (SP), with fixed-order RCCL tensor
transport.

Note

Validation is operator-only. It uses seeded tensors and does not load or
benchmark a model, checkpoint, tokenizer, dataset, or serving engine.

Comparison contract

The three experiments are intentionally independent:

Question Baseline Metric
Does distributed execution preserve determinism? This PR's deterministic Triton FFN at TP=1 Element mismatch count for forward output, training output, dHidden, and dWeights; acceptance is 0
What is the deterministic performance cost? Official Hugging Face Transformers Qwen3MLP at TP=1 Median FFN latency only; no accuracy comparison is mixed into the speed result
What is the simple FP16 precision observation? Official Qwen3MLP at TP=1 in FP32 FP16 output relative-L2, max-absolute, and mean-absolute error

Native ROCm/RCCL is not used as a numerical-accuracy reference in this report.

Design

  • Implement the bias-free gated Qwen3 FFN directly with ROCm-native Triton
    kernels; there is no CUDA-generated HIP source in this PR.
  • Use a canonical FP32-leaf/BF16-node midpoint reduction tree in deterministic
    GEMM and preserve BF16 stage boundaries in forward and backward.
  • Make each contiguous TP K shard the same subtree used by TP=1.
  • Use RCCL for fixed rank-order tensor transport, followed by a fixed balanced
    BF16 rank reduction tree.
  • Gather complete CP token sequences before weight-gradient GEMMs so their K
    tree matches CP=1.
  • Support TP all-reduce, SP all-gather/reduce-scatter, CP all-gather, and all
    corresponding backward paths.

Operator test matrix

Weights use Hugging Face [out, in] layout. No model-level benchmark and no
separate gate/up/down projection benchmark is included.

Experiment Shape / dtype Parallel configurations
Single-GPU FFN speed (M,H,I)=(1/8/32,4096,12288), BF16 Triton TP1 vs official Qwen3MLP TP1; forward and forward+backward
Distributed FFN speed (M,H,I)=(32,4096,12288), BF16 TP2, TP2+SP, TP4, TP2+CP2, TP2+CP2+SP, TP8, TP4+CP2, TP4+CP2+SP; every row vs official TP1
Distributed exactness Same full logical input and weights, BF16 Every TP/CP/SP layout vs deterministic Triton TP1 exact slices
FP16/FP32 observation (M,H,I)=(8,4096,12288) Official Qwen3MLP TP1 FP16 vs the same operator in FP32

ROCm environment

Item Value
GPU 8 × AMD Instinct MI300X
Architecture gfx942
PyTorch 2.12.0+rocm7.14.0a20260608
ROCm runtime 7.14.60850
Transformers 5.10.4
Benchmark implementation commit 08f47d97d0443c5998b8da6b41a22fdf3848da8f
Result and figure commit e64abab

Correctness results

Validation Result
Single-GPU FFN plus real RCCL TP/CP/SP topology suite 17 passed
Formal TP/CP/SP forward output vs Triton TP1 0 mismatch
Formal TP/CP/SP training output vs Triton TP1 0 mismatch
Formal TP/CP/SP dHidden vs Triton TP1 0 mismatch
Formal TP/CP/SP sharded dWeights vs Triton TP1 0 mismatch
Repeated execution and training/inference forward 0 mismatch

Commands used:

NCCL_IB_DISABLE=1 pytest -q \
  tests/test_qwen_ffn.py \
  tests/distributed/test_qwen_ffn_topology.py

NCCL_IB_DISABLE=1 python benchmarks/benchmark_rocm_ffn.py \
  --warmup 3 \
  --samples 10 \
  --training-samples 5 \
  --output-dir benchmarks/results/pr325_rocm_mi300x

Performance results

All speed numbers compare complete FFN calls. The official baseline is upstream
Transformers Qwen3MLP with unsharded weights and input at TP=1. Distributed
timing uses synchronized wall time and the slowest rank per sample. These rows
compare speed only.

Scope Deterministic Triton / official Qwen3MLP TP1 median latency
Single GPU, forward, M=1/8/32 9.03-22.86x
Single GPU, forward+backward, M=1/8/32 7.38-11.56x
Distributed, forward, eight TP/CP/SP layouts 8.76-15.79x
Distributed, forward+backward, eight TP/CP/SP layouts 7.45-14.06x

The separate dtype observation runs only official Qwen3MLP TP1:

Candidate Reference Relative L2 Max abs Mean abs
FP16 FP32 6.544e-4 (0.06544%) 2.046e-6 3.742e-7

Full combined report ·
Raw JSON

Single-GPU official TP1 versus Triton speed

Topology mismatch versus Triton TP1

Distributed official TP1 versus Triton speed

Communication overlap assessment

The reported implementation serializes dependencies; the measurements do not
claim communication/computation overlap.

  • Forward SP all-gather must complete before gate/up computation, and the final
    TP reduction consumes the down-projection output. These are hard dependencies.
  • In backward, the gate and up contributions to dHidden are independent until
    their final ordered addition. A future implementation can reduce one on a
    second stream while computing the other.
  • That optimization must preserve rank order, the reduction tree, wait points,
    BF16 stage boundaries, and gate-then-up addition order. It is accepted only if
    every TP1 mismatch column remains zero.

Communication implementation provenance

The ROCm deterministic communication operator used by this PR is adopted from PR #357. The current path uses the optimized HIP IPC fixed-tree implementation with RCCL fallback, including the packed reduce_scatter_many path for the independent sequence-parallel FFN backward lanes.

The checked benchmark report records this implementation provenance and keeps the main performance figure as a four-way same-topology comparison without adding a PR-specific series.

@coderabbitai

coderabbitai Bot commented Aug 20, 2026

Copy link
Copy Markdown

Important

Review skipped

Auto reviews are disabled on base/target branches other than the default branch.

Please check the settings in the CodeRabbit UI or the .coderabbit.yaml file in this repository. To trigger a single review, invoke the @coderabbitai review command.

⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro Plus

Run ID: 9aa2becf-9e06-4ed0-8a5f-2e3d9cb4b5fe

You can disable this status message by setting the reviews.review_status to false in the CodeRabbit configuration file.

Use the checkbox below for a quick retry:

  • 🔍 Trigger review

Comment @coderabbitai help to get the list of available commands.

@Flink-ddd
Flink-ddd changed the base branch from codex/ws2-rocm-strict-attention to main August 21, 2026 15:37
@frank-2077
frank-2077 changed the base branch from main to codex/ws2-rocm-strict-attention August 21, 2026 15:40
zhangj1an and others added 29 commits August 29, 2026 08:38
The previous commit renamed the reference path reference-hip -> reference-native in
the code (the row is the CUDA build of the same .cu on an NVIDIA host, so the old
name would be wrong in an H100 report) but left the recorded results.json and
report.md on the old key. Nothing errored: PATH_NAMES simply stopped matching, so
regenerating the report silently dropped the reference row entirely -- the one row
the Triton core's whole bitwise claim is measured against.

Rename the key in the stored single-GPU cases and batch-composition rows, record the
production_path the strict-vs-reference gap was measured against, and regenerate.
Values are untouched; this is the same kernel under a name that is honest on both
platforms.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01824WkgMDxtBKD4Ex2NkfYA
…e script

Three fixes plus the CPU numbers.

The per-KV-group schedule cost was measured by a throwaway script and merged into
results.json by hand, so the full MI300X re-run silently dropped it. It is now
_tp_schedule_cost() inside the benchmark and part of the normal flow. Re-measured:
3.60-7.95x the single-launch forward (previously reported 4.11-7.31x; same
conclusion, run-to-run variance).

_environment() reported device facts regardless of device, so the host column
claimed gpu_count=8, hip=7.14 and an RCCL collective for a run that never touched a
GPU. It now zeroes those on a host run.

Host results, S<=2048, BF16: sdpa 31.5/98.3/304.6 ms and pytorch-native
15.4/51.6/216.2 ms forward at S=512/1024/2048. Only those two paths exist on the
host -- strict-aiter is ROCm-only and the reference and Triton cores need a GPU --
which is the same shape as PR RL-Align#328's CPU column. S=4096 was dropped after a first
attempt was killed at 25 minutes; the host column is absolute-latency context, not
a headline.

MI300X was re-run so every platform has pytorch-native, the common path, and so the
reference core carries its platform-neutral name.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01824WkgMDxtBKD4Ex2NkfYA
The CPU sweep at S=4096 could not be priced in advance. My S^2 estimate said 24
minutes; ten minutes in it was still inside the first path, because the materialized
score matrix (4.3-8.6 GB) leaves cache and the run becomes memory-bandwidth bound,
so the compute model does not hold at that size. Two runs were killed by hand on
guesswork.

Price each (path, case) from the one untimed call that already runs to capture the
outputs, project the sampling cost from it, and skip the cell when that exceeds
--path-budget-seconds (default 900). A skipped cell is reported, not dropped: the
report grows a "Skipped cells" section carrying the observed single-call cost and
the projection, and the latency row reads "skipped" rather than going blank. The
FP64 accuracy figure survives, since it only needs the one call.

The observed cost is the useful part. fp16 pytorch-native at S=2048 is 13.2 s per
forward against 216 ms for the same shape in bf16 -- a 61x gap that measures
PyTorch's missing fp16 CPU matmul, not this operator.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01824WkgMDxtBKD4Ex2NkfYA
Signed-off-by: vensen <vensenmu@gmail.com>
CPU host numbers for the two paths that exist without a device: torch SDPA and
NativeAttentionOp. S=512/1024/2048, bf16 and fp16. Timing is wall clock and peak
memory an RSS high-water delta from /proc, so the host figures approximate and are
not directly comparable to the device columns.

The budget is now enforced twice: a pre-flight projection from the one untimed call
that already runs to capture outputs, and a wall clock inside the sampling loops,
because the projection under-estimates once the materialized score matrix leaves
cache. A cell that runs out of budget is truncated and flagged rather than dropped,
and one that cannot start is listed under "Skipped cells" with its observed
single-call cost.

Worth knowing before reading the host column: fp16 pytorch-native is 21.1 s per
forward at S=2048 against 228 ms for the same shape in bf16. That is PyTorch having
no optimized fp16 CPU matmul, not a property of this operator.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01824WkgMDxtBKD4Ex2NkfYA
Signed-off-by: maxiaosong1124 <maxiaosong7890@outlook.com>
Port the HIP IPC fixed-tree transport and packed reduce-scatter path from PR RL-Align#357, with RCCL fallback and focused ROCm coverage.
Record the PR RL-Align#357 collective measurements, preserve the four-way same-topology comparison, and publish the refreshed MI300X artifacts.
feat(cuda): promote deterministic cross-config runtime and kernel validation to main
# Conflicts:
#	csrc/ops.cpp
#	rl_engine/distributed/collectives.py
#	rl_engine/kernels/ops/pytorch/ffn/ffn.py
…tch path

CP orchestration lived only in StrictCUDAAttentionRuntime, so ROCm had nowhere
to put it: the AITER/CK core is single-rank arithmetic, the Vime provider failed
closed at CP>1, and the only working AG/core/RS sequence was in the benchmark
script. StrictRocmAttentionRuntime mirrors the CUDA runtime over the RCCL AG/RS
transport.

Two things differ from the CUDA runtime and both are load-bearing:

- The core is launched once per (batch row, KV group) rather than once per
  sequence. AITER/CK's reduction order depends on how many heads shared the
  launch, so a head shard computed under TP=N is otherwise not bit-identical to
  the same shard under a different TP degree. FA4 has no such dependence.
- RCCL moves tensors but never reduces them. The cross-rank (out, lse) combine
  is the transport's fixed balanced rank tree, not RCCL's own algorithm
  selection, which varies with message size and topology.

The sequence reorder and position validation are bound from the CUDA runtime
rather than reimplemented, so the two runtimes cannot drift into two different
global orderings. The per-KV-group launch loop moves out of the Vime provider
into the runtime, so CP=1 and CP>1 now share one schedule instead of keeping a
second copy in the provider.

Opening CP also required the registry to stop rejecting it: cp_world_sizes was
(1,) and deterministic_cp_merge was False, so AttentionBackendCapability
rejected CP>1 twice over. cp_world_sizes now matches the world sizes the RCCL
transport accepts and deterministic_cp_merge is True because the merge order is
ours. A test pins the two together so the declaration cannot drift from the
transport. zigzag fails closed: the strict CP plan describes one contiguous
block per rank, and a zigzag rank owns two discontiguous runs.

Measured on 8xMI300X through attention_provider, not the transport directly, so
the test also pins that CP is reachable from the production dispatch path:
CP=2/4/8 are bitwise against a CP=1 run of the same core on the same logical
sequence, 0 mismatched elements on out and lse, repeat-bitwise on every rank.

Also corrects the stale TP comment the PR description flagged in section 7: the
shipped policy removes the degree dependence with the per-KV-group launch rather
than binding the degree to avoid a ~3x cost.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01NPU7GosGZSdj6pKBWX7n2Y
… collective

The ROCm attention CP path owned a second transport implementation.
RCCLAGRSAttentionCPCommunication overrode _get_collective to construct
_RCCLRankOrderedTransport, which bypassed the collective_for_group factory
that the CUDA adapter goes through and that already dispatches to
RCCLDeterministicCollective on HIP.

Both copies evaluated the same balanced rank tree, so the two platforms were
bit-identical -- but only by coincidence. Nothing pinned them together, so a
later change to the shared collective's reduction order would have left the
attention path on the old tree with no test failing.

Delete _RCCLRankOrderedTransport and the override. ROCm now inherits the CUDA
adapter's resolution, so one implementation serves both platforms. The
reduction expression is unchanged, so this is not expected to move any bit.

_RootReduceScatterSequence falls back from scatter() to reduce_scatter()
because the shared collective exposes no scatter entrypoint. That is the same
branch the CUDA path has always taken and it is semantically equivalent --
non-root ranks zero their input, so the tree sum returns the root's value and
adding zero is exact. It costs one extra all-gather plus tree per call.

The error messages the adapters raise are now keyed off a collective_label
class attribute so the inherited path still reports RCCL on ROCm instead of
mislabeling itself as CUDA.

Three tests pin the arrangement: the two adapters must share one
_get_collective, the ROCm adapter must resolve through collective_for_group,
and the registry's cp_world_sizes must equal the shared collective's
_SUPPORTED_WORLD_SIZES so the capability declaration cannot drift from what
the transport accepts.

Not verified on MI300X. The reduction expression is unchanged, but the
reduce_scatter fallback and the shared collective's capacity and signature
validation are new to this path, so a CP=2/4/8 bitwise run is still owed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: zhangj1an <jianmusings@gmail.com>
Signed-off-by: vensen <vensenmu@gmail.com>
Signed-off-by: vensen <vensenmu@gmail.com>
Signed-off-by: maxiaosong1124 <maxiaosong7890@outlook.com>
…codex/ws2-rocm-strict-attention

Brings in the ROCm HIP IPC deterministic collective so ROCm gets a native
transport instead of the RCCL-only Python path. Because this branch already
routes the attention CP adapter through collective_for_group, the ROCm
attention CP path now resolves to that HIP IPC collective with no further
change.

Conflict resolutions:

* cp_comm.py -- PR RL-Align#357 optimizes _RCCLRankOrderedTransport.scatter; this
  branch deleted that class in favour of the shared collective. Kept the
  deletion: the optimization targets code that no longer exists, and the
  shared collective supersedes it.
* collectives.py -- kept both sides' module constants (they are additive:
  CUDA staging-frame sizes and ROCm IPC tuning thresholds). reduce_scatter_many
  had diverged signatures, so the merged one takes the union: PR RL-Align#357's
  inputs/outs plus this branch's validate_signature, forwarding both.
* ffn.py -- PR RL-Align#357 restructured the backward to compute both gate and up
  input gradients up front and pack them into one reduce_scatter_many, while
  this branch renamed _gemm_fwd to _linear_da/_linear_dw. Took PR RL-Align#357's
  structure with this branch's helper names; the old second reduce-scatter
  for the up lane is gone.
* setup.py -- kept this branch's Ascend build imports (sysconfig,
  CompileError, find_executable, Extension are used further down the file)
  and added PR RL-Align#357's ROCm .hip source to cuda_sources.
* ops.cpp, _C.pyi -- both additive; kept both sides.

tests/distributed and tests/test_build_platform_collectives: 46 passed,
5 skipped. tests/test_qwen_ffn.py fails 24 here, but 25 of the same tests
fail on the pre-merge tree in this environment: the extension was built
without KERNEL_ALIGN_DET_GEMM_SM90=1, so strict GEMM refuses to run. The
merge removes one of those failures and adds none.

Not verified on MI300X. The HIP IPC path has no coverage in this
environment, so the ROCm CP bitwise run is still owed.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: zhangj1an <jianmusings@gmail.com>
… into

The CUDA core checks the FA4 CuTe API by parameter name before it runs, so a
renamed or dropped strict control fails at load. ROCm had no equivalent:
inspect.signature reports (*args, **kwargs) for AITER's JIT wrapper, so the
only guard was a SHA-256 of the module source. That catches "something
changed" but cannot say what, and it fires on unrelated edits.

Read the registered Torch schema instead (torch.ops.aiter.<op>.default).

The check is an ordered prefix, not a name set, because the two call sites
pass positionally. An argument inserted upstream would shift the meaning of
every later argument while the call still type-checks -- dropout_p, the two
window sizes and sink_size are all int/bool, so nothing would raise and the
kernel would run with silently reinterpreted controls. Name presence alone
does not catch that; the FA4 path is exempt only because it calls by keyword.

This also pins something that was previously unprovable: the True at
backward position 11 is the schema's `deterministic`. ROCm's backward was
already deterministic, but nothing tied that literal to its parameter.

The source fingerprint stays as a second line: the schema check describes what
changed, the fingerprint still catches a same-schema implementation change.

Verified against the installed AITER; the forward prefix matches exactly.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: zhangj1an <jianmusings@gmail.com>
StrictCUDAAttentionRuntime has forward_paged_with_lse; the ROCm runtime had
no decode entry point at all, so decode-stage KV-cache replay existed only in
the device-neutral comparison harness with no ROCm path behind it.

No AITER paged kernel can serve this contract. paged_attention_rocm, _ragged
and _v1 are all two-pass partition reducers -- partition_size, with
exp_sums/max_logits/tmp_out partials -- so the partition count tracks the
cached length and Split-KV cannot be turned off. AITER's
flash_attn_varlen_func takes a block_table but exposes no num_splits to pin,
unlike CUDA's FA4. Either way the contract could not prove Split-KV disabled,
which attention_binding checks from both runtime evidence and the contract.

So the pages are gathered into logical order and handed to the same dense
core the prefill path uses, at the same one-launch-per (batch row, KV group)
granularity. The arithmetic is then identical to a CP=1 prefill over the same
logical sequence, which is what makes decode replay comparable against it. The
cost is materializing the cached KV; a native paged kernel avoids that and can
replace this once AITER can pin its split count.

The registry deliberately does NOT gain AttentionMode.DECODE. Nothing routes
to the new entry point: the Vime provider always calls forward_with_lse and
builds its contract with kv_cache=None, and its request carries no page table.
Declaring the mode now would let the binding layer accept a decode path that
never executes. A test pins the omission so it flips together with the
dispatch wiring rather than drifting ahead of it.

Tests inject the core, so they run without ROCm. They pin the part that is
ours rather than AITER's: a shuffled page table still yields logical KV order,
the gather truncates to seqused_k instead of exposing the page tail, each
launch still sees exactly one KV group, and the provenance says
paged_kernel=none so no reader mistakes this for a native paged path.

Not verified on MI300X. The core arithmetic is unchanged, but the gather's
index_select/reshape/permute and the claimed bitwise equality with a CP=1
prefill over the same tokens both need a real run.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: zhangj1an <jianmusings@gmail.com>
Added a new function to precompile strict attention training for better performance when using RL_KERNEL. Modified the integration entry point to include this precompilation step based on the plan's implementation.
# Conflicts:
#	csrc/ops.cpp
#	rl_engine/distributed/collectives.py
#	rl_engine/kernels/attention_contract.py
#	rl_engine/kernels/ops/cuda/attention/__init__.py
#	rl_engine/kernels/ops/cuda/attention/deterministic_attn.py
#	rl_engine/kernels/ops/cuda/attention/flash_attn.py
#	rl_engine/kernels/ops/pytorch/attention/cp_attention.py
#	scripts/ws2_p2p_nccl_attention_reference_check.py
#	tests/test_det_gemm.py
#	tests/test_flashinfer_pr7_attention.py
#	tests/test_qwen_ffn.py
@frank-2077
frank-2077 changed the base branch from codex/ws2-rocm-strict-attention to test September 1, 2026 03:36
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

platform: rocm Specific tasks specific to AMD graphics cards (such as CK, bpreshuffle/FA)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants