From b509dc6a437bf4b5f0b28c5787ff27dcc9c86def Mon Sep 17 00:00:00 2001 From: hallerite Date: Fri, 24 Apr 2026 19:10:16 +0530 Subject: [PATCH 1/2] fix(ring_attn): adapt FA3 varlen wrappers to flash_attn_3 3.0.0 kwargs MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit flash_attn_3 3.0.0 (released on the pytorch-cu128-test index, picked up in PR #2234 on 2026-04-09) renamed the kernel's `causal` argument to `is_causal` and split `window_size` into `window_size_left` / `window_size_right`. The FA3 ring-attn wrapper here still updates the params dict with the old `causal` key, which leaves the new `is_causal` default in place and passes both — triggering: RuntimeError: flash_attn_3::_flash_attn_backward() expected at most 22 argument(s) but received 23 on the first backward pass for any model+config that routes through `_fa3_varlen_forward` / `_fa3_varlen_backward` (cp > 1 with FA3, custom impl). Forward hit it first in our Qwen3.5 run; backward has the same pattern. Fix: pick the correct key based on what the installed kernel's signature exposes. Supports both the newer 3.0.0 naming and the older 3.0.0b1 naming without pinning the kernel version. Validated against flash_attn_3 3.0.0 (pytorch-cu128-test) by running a Qwen/Qwen3.5-35B-A3B RL config with cp=2 + FA3 + custom impl — the forward pass now succeeds end-to-end (loss computed); backward uses the same kwarg-passing pattern. --- .../trainer/models/layers/ring_attn.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/src/prime_rl/trainer/models/layers/ring_attn.py b/src/prime_rl/trainer/models/layers/ring_attn.py index 0db4b53c87..db9b49ddc4 100644 --- a/src/prime_rl/trainer/models/layers/ring_attn.py +++ b/src/prime_rl/trainer/models/layers/ring_attn.py @@ -30,11 +30,19 @@ def _fa3_varlen_forward( "max_seqlen_q": max_seqlen_q, "max_seqlen_k": max_seqlen_k, "softmax_scale": softmax_scale, - "causal": causal, } ) + # flash_attn_3 3.0.0 renamed `causal` -> `is_causal`. Older 3.0.0b1 + # builds kept the `causal` name, so check which key the kernel exposes. + if "is_causal" in params: + params["is_causal"] = causal + else: + params["causal"] = causal if "window_size" in params: params["window_size"] = window_size + elif "window_size_left" in params and "window_size_right" in params: + params["window_size_left"] = window_size[0] + params["window_size_right"] = window_size[1] out, lse, _, _ = _flash_attn_forward(**params) return out, lse @@ -76,11 +84,17 @@ def _fa3_varlen_backward( "dk": dk, "dv": dv, "softmax_scale": softmax_scale, - "causal": causal, } ) + if "is_causal" in params: + params["is_causal"] = causal + else: + params["causal"] = causal if "window_size" in params: params["window_size"] = window_size + elif "window_size_left" in params and "window_size_right" in params: + params["window_size_left"] = window_size[0] + params["window_size_right"] = window_size[1] _flash_attn_backward(**params) From 66b277fb02b24b9d4be9846cfaaa50aaf1d38af7 Mon Sep 17 00:00:00 2001 From: hallerite Date: Fri, 24 Apr 2026 19:23:38 +0530 Subject: [PATCH 2/2] refactor(ring_attn): drop flash_attn_3 version fallback The repo pins flash_attn_3 via pytorch-cu128-test (3.0.0), and uv.lock resolves exactly that. The older 3.0.0b1 kernel is unreachable through any supported install path, so keeping conditional branches for its `causal` / `window_size` argument names is dead code. Collapse to the single 3.0.0 API: `is_causal` + `window_size_left` + `window_size_right` set unconditionally. One way to do it. --- .../trainer/models/layers/ring_attn.py | 26 +++++-------------- 1 file changed, 6 insertions(+), 20 deletions(-) diff --git a/src/prime_rl/trainer/models/layers/ring_attn.py b/src/prime_rl/trainer/models/layers/ring_attn.py index db9b49ddc4..34e30e41c6 100644 --- a/src/prime_rl/trainer/models/layers/ring_attn.py +++ b/src/prime_rl/trainer/models/layers/ring_attn.py @@ -30,19 +30,11 @@ def _fa3_varlen_forward( "max_seqlen_q": max_seqlen_q, "max_seqlen_k": max_seqlen_k, "softmax_scale": softmax_scale, + "is_causal": causal, + "window_size_left": window_size[0], + "window_size_right": window_size[1], } ) - # flash_attn_3 3.0.0 renamed `causal` -> `is_causal`. Older 3.0.0b1 - # builds kept the `causal` name, so check which key the kernel exposes. - if "is_causal" in params: - params["is_causal"] = causal - else: - params["causal"] = causal - if "window_size" in params: - params["window_size"] = window_size - elif "window_size_left" in params and "window_size_right" in params: - params["window_size_left"] = window_size[0] - params["window_size_right"] = window_size[1] out, lse, _, _ = _flash_attn_forward(**params) return out, lse @@ -84,17 +76,11 @@ def _fa3_varlen_backward( "dk": dk, "dv": dv, "softmax_scale": softmax_scale, + "is_causal": causal, + "window_size_left": window_size[0], + "window_size_right": window_size[1], } ) - if "is_causal" in params: - params["is_causal"] = causal - else: - params["causal"] = causal - if "window_size" in params: - params["window_size"] = window_size - elif "window_size_left" in params and "window_size_right" in params: - params["window_size_left"] = window_size[0] - params["window_size_right"] = window_size[1] _flash_attn_backward(**params)