diff --git a/unsloth/models/cohere.py b/unsloth/models/cohere.py
index 0b7f3ab973..cb367d451e 100644
--- a/unsloth/models/cohere.py
+++ b/unsloth/models/cohere.py
@@ -22,6 +22,7 @@ from ..utils.attention_dispatch import (
AttentionContext,
run_attention,
select_attention_backend,
+ resolve_prefix_seg_info,
)
try:
@@ -151,6 +152,9 @@ def CohereAttention_fast_forward(
"softmax_scale": getattr(self, "softmax_scale", None),
},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse
+ # (KV cache / padding mask) raises. None => byte-identical default.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -161,6 +165,7 @@ def CohereAttention_fast_forward(
seq_info = seq_info,
attention_mask = attention_mask,
causal_mask = causal_mask,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V)
diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py
index 68b9ebe22f..4a0531db78 100644
--- a/unsloth/models/gemma2.py
+++ b/unsloth/models/gemma2.py
@@ -22,6 +22,7 @@ from ..utils.attention_dispatch import (
AttentionContext,
run_attention,
select_attention_backend,
+ resolve_prefix_seg_info,
SDPA,
)
from .gemma import (
@@ -168,6 +169,11 @@ def Gemma2Attention_fast_forward(
},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse
+ # (KV cache / padding mask) raises. None => byte-identical default. gemma2 is
+ # sliding-window and softcapped: the engage gate caps spans at the window and
+ # excludes softcap models entirely, so PG never engages here.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -179,6 +185,7 @@ def Gemma2Attention_fast_forward(
attention_mask = attention_mask,
causal_mask = causal_mask,
sliding_window = sliding_window,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V)
diff --git a/unsloth/models/granite.py b/unsloth/models/granite.py
index f5b0f57aa6..4dedf642eb 100644
--- a/unsloth/models/granite.py
+++ b/unsloth/models/granite.py
@@ -23,6 +23,7 @@ from ..utils.attention_dispatch import (
AttentionContext,
run_attention,
select_attention_backend,
+ resolve_prefix_seg_info,
SDPA,
)
from .llama import (
@@ -159,6 +160,9 @@ def GraniteAttention_fast_forward(
},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse
+ # (KV cache / padding mask) raises. None => byte-identical default.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -169,6 +173,7 @@ def GraniteAttention_fast_forward(
seq_info = seq_info,
attention_mask = attention_mask,
causal_mask = causal_mask,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V)
diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py
index 6bec95b577..417a88f480 100644
--- a/unsloth/models/llama.py
+++ b/unsloth/models/llama.py
@@ -39,6 +39,7 @@ from ..utils.attention_dispatch import (
run_attention,
SDPA,
select_attention_backend,
+ resolve_prefix_seg_info,
)
from torch.nn.functional import scaled_dot_product_attention
from transformers import __version__ as transformers_version
@@ -738,6 +739,10 @@ def LlamaAttention_fast_forward(
flash_dense_kwargs = {"causal": True},
flash_varlen_kwargs = {"dropout_p": 0.0, "causal": True},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward (same route
+ # as packed_seq_lengths); misuse (KV cache / padding mask) raises. None => byte-identical
+ # default. Reuse of this forward also carries the branch to qwen2 & gemma.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -748,6 +753,7 @@ def LlamaAttention_fast_forward(
seq_info = seq_info,
attention_mask = attention_mask,
causal_mask = causal_mask,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = config, context = context, Q = Q, K = K, V = V)
@@ -895,8 +901,10 @@ def LlamaModel_fast_forward(
seq_length_with_past = seq_length
# Fix out of bounds tokenization unless we were given packed metadata
- allow_overlength = getattr(self, "_unsloth_allow_packed_overlength", False) or (
- "packed_seq_lengths" in kwargs
+ allow_overlength = (
+ getattr(self, "_unsloth_allow_packed_overlength", False)
+ or ("packed_seq_lengths" in kwargs)
+ or ("prefix_seg_info" in kwargs and kwargs["prefix_seg_info"] is not None)
)
if hasattr(self, "max_seq_length") and not allow_overlength:
if seq_length > self.max_seq_length:
diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py
index df2a4de5bd..4350565fe2 100644
--- a/unsloth/models/mistral.py
+++ b/unsloth/models/mistral.py
@@ -27,6 +27,7 @@ from ..utils.attention_dispatch import (
run_attention,
SDPA,
select_attention_backend,
+ resolve_prefix_seg_info,
)
from .llama import (
LlamaRotaryEmbedding,
@@ -124,6 +125,9 @@ def MistralAttention_fast_forward(
"softmax_scale": getattr(self, "softmax_scale", None),
},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse
+ # (KV cache / padding mask) raises. None => byte-identical default.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -134,6 +138,7 @@ def MistralAttention_fast_forward(
seq_info = seq_info,
attention_mask = attention_mask,
causal_mask = causal_mask,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V)
@@ -161,7 +166,13 @@ def MistralForCausalLM_fast_forward(
*args,
**kwargs,
) -> Union[Tuple, CausalLMOutputWithPast]:
- if causal_mask is None and past_key_values is None:
+ # PrefixGrouper brings its own mask: a synthesized causal attention_mask would trip
+ # resolve_prefix_seg_info on the no-xFormers path and force a fallback.
+ if (
+ causal_mask is None
+ and past_key_values is None
+ and kwargs.get("prefix_seg_info", None) is None
+ ):
bsz, q_len = input_ids.shape
sliding_window = getattr(self.config, "sliding_window", None)
diff --git a/unsloth/models/qwen3.py b/unsloth/models/qwen3.py
index e28e72d3ea..0d05a2d538 100644
--- a/unsloth/models/qwen3.py
+++ b/unsloth/models/qwen3.py
@@ -23,6 +23,7 @@ from ..utils.attention_dispatch import (
run_attention,
SDPA,
select_attention_backend,
+ resolve_prefix_seg_info,
)
from .llama import (
LlamaRotaryEmbedding,
@@ -146,6 +147,9 @@ def Qwen3Attention_fast_forward(
"softmax_scale": getattr(self, "softmax_scale", None),
},
)
+ # PrefixGrouper seg table rides in **kwargs from the GRPO logprob forward; misuse
+ # (KV cache / padding mask) raises. None => byte-identical default.
+ _pg_seg = resolve_prefix_seg_info(kwargs, past_key_value, attention_mask)
context = AttentionContext(
bsz = bsz,
q_len = q_len,
@@ -156,6 +160,7 @@ def Qwen3Attention_fast_forward(
seq_info = seq_info,
attention_mask = attention_mask,
causal_mask = causal_mask,
+ prefix_seg_info = _pg_seg,
)
A = run_attention(config = attention_config, context = context, Q = Q, K = K, V = V)
diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py
index 0573fd5fd8..098950de08 100644
--- a/unsloth/models/rl_replacements.py
+++ b/unsloth/models/rl_replacements.py
@@ -31,6 +31,7 @@ from unsloth_zoo.rl_replacements import (
left_pack_padding,
create_completion_attention_mask,
chunked_selective_log_softmax,
+ chunked_hidden_states_selective_log_softmax,
_unsloth_get_mm_token_id,
_unsloth_fix_mm_token_type_ids,
)
@@ -65,6 +66,25 @@ try:
)
except Exception:
UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = False
+# One-time PrefixGrouper gate; any import failure degrades to "PrefixGrouper off".
+_pg_build_layout = _pg_enabled_fn = _pg_verify_on = _pg_tol_ok = _PG_TOL_KILL = None
+UNSLOTH_GRPO_PREFIX_GROUPER_ON = os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER", "1").lower() not in (
+ "0",
+ "false",
+ "no",
+ "off",
+)
+if UNSLOTH_GRPO_PREFIX_GROUPER_ON:
+ try:
+ from ..utils.prefix_grouper import (
+ build_group_layout as _pg_build_layout,
+ prefix_grouper_enabled as _pg_enabled_fn,
+ verify_on as _pg_verify_on,
+ tol_ok as _pg_tol_ok,
+ TOL_KILL as _PG_TOL_KILL,
+ )
+ except Exception:
+ UNSLOTH_GRPO_PREFIX_GROUPER_ON = False
RL_EXTRA_ARGS = defaultdict(list)
RL_FUNCTIONS = defaultdict(list)
@@ -1380,6 +1400,166 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
# left-pad RoPE error). Self-verified against the per-row forward, re-checked as T
# grows; falls back if a backend ignores packed_seq_lengths.
logprobs = None
+
+ # ---- PrefixGrouper (GRPO shared-prompt dedup; default ON, exact + self-verified) ----
+ # G completions per prompt share the prefix; the packed path forwards it G times,
+ # PrefixGrouper stores it once (FlexAttention shared-prefix mask), cutting the trunk
+ # forward from G*(P+R) to P+G*R tokens. Gated by UNSLOTH_GRPO_PREFIX_GROUPER (needs
+ # seq-packing), tok_r auto-gate, and first-use self-verify vs the packed path
+ # (mismatch => fall back + mark unsafe), so a mask/isolation regression cannot ship
+ # silently. When off / ungrouped / unverified, the packed path below runs as before.
+ _pg_result = None
+ _pg_use = False
+ _pg_skip_pk = False # once a shape is PG-verified, skip the full-row forward
+ _pg_forward_fn = None # deferred PG forward (runs at the verify site below)
+ _pg_num_gen = getattr(self, "num_generations", None)
+ # Env gate hoisted to module level (mirrored via RL_PRE_ITEMS). Skip PG under vLLM
+ # (fast_inference=True): the rollout dominates the step, so PG saves little and its
+ # first-use self-verify is net overhead.
+ _pg_engage = (
+ UNSLOTH_GRPO_PREFIX_GROUPER_ON
+ and not getattr(self, "use_vllm", False)
+ and not getattr(unwrapped_model, "_unsloth_prefix_grouper_nograd_disabled", False)
+ )
+ if _pg_engage:
+ try:
+ # Skip softcap models (the flex kernel never applies attn_logit_softcapping)
+ # and hybrid SSM / MoE models: only the threaded attention forwards get the
+ # shared-prefix isolation, so a Mamba or MoE decoder that does not forward
+ # prefix_seg_info would leak suffixes across completions. PG also rides on
+ # sequence packing, so it needs the same zoo masked-column guard.
+ _pg_cfg = getattr(unwrapped_model, "config", None)
+ _pg_engage = (
+ _pg_enabled_fn()
+ and UNSLOTH_ZOO_HAS_MASKED_COL_GUARD
+ and pixel_values is None
+ and token_type_ids is None
+ and mm_token_type_ids is None
+ and _pg_num_gen is not None
+ and _pg_num_gen >= 2
+ and not getattr(_pg_cfg, "attn_logit_softcapping", None)
+ # normal backends apply config.attention_dropout in training; the flex
+ # path is deterministic, so skip PG when it is set.
+ and not getattr(_pg_cfg, "attention_dropout", 0)
+ and not any(
+ getattr(_pg_cfg, _pg_a, None) is not None
+ for _pg_a in (
+ "mamba_d_ssm",
+ "mamba_d_state",
+ "mamba_expand",
+ "num_experts",
+ "num_local_experts",
+ "n_routed_experts",
+ "moe_intermediate_size",
+ )
+ )
+ )
+ except Exception:
+ _pg_engage = False
+ if _pg_engage:
+ try:
+ _pg_pad = self.processing_class.pad_token_id
+ # cap the PG span (P+max(R)) at the sliding window, like the packed _pk_sw guard.
+ _pg_sw = getattr(
+ getattr(unwrapped_model, "config", None), "sliding_window", None
+ )
+ if not (isinstance(_pg_sw, int) and _pg_sw > 0):
+ _pg_sw = None
+ _pg_layout = _pg_build_layout(
+ input_ids,
+ logits_to_keep,
+ _pg_pad,
+ _pg_num_gen,
+ left_pad_tokens_per_prompt,
+ max_segment_cap = _pg_sw,
+ )
+ _pg_unsafe = getattr(
+ unwrapped_model, "_unsloth_prefix_grouper_nograd_unsafe", None
+ )
+ if _pg_unsafe is None:
+ _pg_unsafe = set()
+ if _pg_layout is not None and _pg_layout.signature not in _pg_unsafe:
+ _pg_sig = _pg_layout.signature
+ _pg_verified = getattr(
+ unwrapped_model, "_unsloth_prefix_grouper_nograd_verified", None
+ )
+ if _pg_verified is None:
+ _pg_verified = set()
+ _pg_chunks = max(1, total_rows * multiplier)
+
+ def _pg_run_forward(_pg_layout = _pg_layout, _pg_chunks = _pg_chunks):
+ with _get_inference_mode_context_manager(model):
+ with torch.amp.autocast(
+ device_type = "cuda", dtype = self._autocast_dtype
+ ):
+ _pg_hidden = unwrapped_model(
+ input_ids = _pg_layout.flat_ids,
+ position_ids = _pg_layout.position_ids,
+ prefix_seg_info = _pg_layout.prefix_seg_info,
+ use_cache = False,
+ ).logits
+ _pg_r = _pg_layout.extract_logps(
+ _pg_hidden,
+ lm_head,
+ chunked_hidden_states_selective_log_softmax,
+ _pg_chunks,
+ logit_scale_multiply,
+ logit_scale_divide,
+ logit_softcapping,
+ temperature,
+ )
+ _pg_hidden = None # release before any verify forward
+ device_synchronize()
+ # clip to the loss window [B, logits_to_keep+max_left_pad]
+ _pg_w = logits_to_keep + max_left_pad
+ if _pg_r.shape[1] > _pg_w:
+ _pg_r = _pg_r[:, -_pg_w:]
+ return _pg_r
+
+ # trust only within the verified envelope: re-verify when T or the
+ # longest segment grows, like the packed path
+ _pg_T = int(_pg_layout.flat_ids.shape[1])
+ _pg_maxseg = int(_pg_layout.position_ids.max()) + 1
+ _pg_env = (
+ _pg_verified.get(_pg_sig) if isinstance(_pg_verified, dict) else None
+ )
+ if (not _pg_verify_on()) or (
+ _pg_env is not None and _pg_T <= _pg_env[0] and _pg_maxseg <= _pg_env[1]
+ ):
+ # trusted shape: run PG now and skip the full-row forward below
+ _pg_result = _pg_run_forward()
+ _pg_use = True
+ _pg_skip_pk = True
+ else:
+ # unverified shape: defer the forward until the packed reference
+ # exists (verify site below), so a declined packed path never wastes
+ # a whole-batch PG forward
+ _pg_forward_fn = _pg_run_forward
+ except Exception as _pg_err:
+ _pg_result = None
+ _pg_use = False
+ _pg_skip_pk = False
+ _pg_forward_fn = None
+ # A FlexAttention/Triton compile failure or OOM here is GPU-wide, not
+ # layout-specific, so retrying the same PG forward every step just re-pays
+ # the failure. Persistently disable PG (mirrors the seq-packing handler
+ # setting _unsloth_seq_packing_nograd_ok = False); the packed/padded path
+ # below still produces the exact result.
+ unwrapped_model._unsloth_prefix_grouper_nograd_disabled = True
+ if isinstance(_pg_err, torch.cuda.OutOfMemoryError):
+ torch.cuda.empty_cache()
+ os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
+ if UNSLOTH_ENABLE_LOGGING:
+ print(
+ f"[Unsloth] GRPO PrefixGrouper (no-grad) disabled (fell back to packed): {_pg_err!r}",
+ flush = True,
+ )
+
+ # ---- Sequence packing (default-on; disable with UNSLOTH_GRPO_SEQ_PACKING=0) ----
+ # One varlen [1, sum L] block-diagonal forward replaces the padded [B, Lmax] loop
+ # (exact per-row result; also fixes the padded path's left-pad RoPE error).
+ # Self-verified vs the per-row forward, re-checked as T grows; falls back if a
+ # backend ignores packed_seq_lengths. lm_head runs on completion positions only.
_pk_result = None
_pk_use = False
_pk_enabled = UNSLOTH_GRPO_SEQ_PACKING_ON
@@ -1388,6 +1568,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
_pk_ok = getattr(unwrapped_model, "_unsloth_seq_packing_nograd_ok", None)
if (
_pk_enabled
+ and not _pg_skip_pk
and pixel_values is None
and token_type_ids is None
and mm_token_type_ids is None
@@ -1462,7 +1643,7 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
)[0]
# GPT-OSS offload race guard (matches the padded loop)
device_synchronize()
- # scatter each completion logprob back to its (row, col) so [:, -_pk_W:] matches padded
+ # scatter each logprob back to its (row, col) so [:, -_pk_W:] matches padded
_pk_tgt = (_pk_nz_idx[1:, 0] * _pk_L + _pk_nz_idx[1:, 1])[_pk_ctgt]
_pk_result = (
torch.zeros(
@@ -1574,7 +1755,73 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
f"[Unsloth] GRPO sequence-packing (no-grad) disabled (fell back to padded): {_pk_err!r}",
flush = True,
)
- if _pk_use and _pk_result is not None:
+ # ---- PrefixGrouper first-use self-verify (no-grad) ----
+ # Compare the untrusted PG result to the full-row packed result (itself verified vs
+ # per-row) over the completion mask: < tol_ok -> trust the structure; >= TOL_KILL ->
+ # unsafe forever; borderline -> fall back this shape.
+ if _pg_forward_fn is not None and not _pg_use:
+ if _pk_use and _pk_result is not None:
+ try:
+ # deferred PG forward, run only now that the packed reference exists
+ _pg_result = _pg_forward_fn()
+ _pg_W2 = logits_to_keep + max_left_pad
+ _pg_cm = create_completion_attention_mask(
+ input_ids[:, -_pg_W2:],
+ left_pad_tokens_per_prompt,
+ max_left_pad,
+ self.processing_class.pad_token_id,
+ ).float()
+ _pg_a = _pg_result[:, -_pg_W2:].float()
+ _pg_b = _pk_result[:, -_pg_W2:].float()
+ _pg_diff = float(((_pg_a - _pg_b).abs() * _pg_cm).max())
+ if UNSLOTH_ENABLE_LOGGING:
+ print(
+ f"[Unsloth] GRPO PrefixGrouper (no-grad) verify: sig={_pg_layout.signature} "
+ f"shared-prefix vs full-row-packed max|d|={_pg_diff:.4f}",
+ flush = True,
+ )
+ if _pg_diff < _pg_tol_ok():
+ _pg_v = getattr(
+ unwrapped_model, "_unsloth_prefix_grouper_nograd_verified", None
+ )
+ if not isinstance(_pg_v, dict):
+ _pg_v = {}
+ _pg_vT = int(_pg_layout.flat_ids.shape[1])
+ _pg_vS = int(_pg_layout.position_ids.max()) + 1
+ _pg_old = _pg_v.get(_pg_layout.signature, (0, 0))
+ _pg_v[_pg_layout.signature] = (
+ max(_pg_vT, _pg_old[0]),
+ max(_pg_vS, _pg_old[1]),
+ )
+ unwrapped_model._unsloth_prefix_grouper_nograd_verified = _pg_v
+ _pg_use = True
+ else:
+ _pg_u = getattr(
+ unwrapped_model, "_unsloth_prefix_grouper_nograd_unsafe", None
+ )
+ if _pg_u is None:
+ _pg_u = set()
+ if _pg_diff >= _PG_TOL_KILL:
+ _pg_u.add(_pg_layout.signature)
+ unwrapped_model._unsloth_prefix_grouper_nograd_unsafe = _pg_u
+ _pg_use = False
+ except Exception as _pg_err3:
+ _pg_result = None
+ _pg_use = False
+ if isinstance(_pg_err3, torch.cuda.OutOfMemoryError):
+ torch.cuda.empty_cache()
+ os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
+ if UNSLOTH_ENABLE_LOGGING:
+ print(
+ f"[Unsloth] GRPO PrefixGrouper (no-grad) verify failed (fell back to packed): {_pg_err3!r}",
+ flush = True,
+ )
+ # else: no packed reference (packing off/failed) -> cannot verify; fall back.
+
+ if _pg_use and _pg_result is not None:
+ logprobs = _pg_result # PrefixGrouper verified/trusted -> skip the loop
+ zipped_inputs = []
+ elif _pk_use and _pk_result is not None:
logprobs = _pk_result # verified -> skip the loop
zipped_inputs = []
else:
@@ -1752,7 +1999,7 @@ RL_PRE_ITEMS["grpo_trainer"].append(
"import os as _unsloth_os\n"
"UNSLOTH_ENABLE_LOGGING = _unsloth_os.environ.get('UNSLOTH_ENABLE_LOGGING', '0') in ('1', 'True', 'true')\n"
)
-# One-time sequence-packing gates, same values as the module-top constants above.
+# Sequence-packing gates, same values as the module-top constants.
RL_PRE_ITEMS["grpo_trainer"].append(
"UNSLOTH_GRPO_SEQ_PACKING_ON = _unsloth_os.environ.get('UNSLOTH_GRPO_SEQ_PACKING', '1').lower() not in ('0', 'false', 'no', 'off')\n"
)
@@ -1764,6 +2011,16 @@ RL_PRE_ITEMS["grpo_trainer"].append(
"except Exception:\n"
" UNSLOTH_ZOO_HAS_MASKED_COL_GUARD = False\n"
)
+# PrefixGrouper gate, same shape as the module-top constants.
+RL_PRE_ITEMS["grpo_trainer"].append(
+ "_pg_build_layout = _pg_enabled_fn = _pg_verify_on = _pg_tol_ok = _PG_TOL_KILL = None\n"
+ "UNSLOTH_GRPO_PREFIX_GROUPER_ON = _unsloth_os.environ.get('UNSLOTH_GRPO_PREFIX_GROUPER', '1').lower() not in ('0', 'false', 'no', 'off')\n"
+ "if UNSLOTH_GRPO_PREFIX_GROUPER_ON:\n"
+ " try:\n"
+ " from unsloth.utils.prefix_grouper import build_group_layout as _pg_build_layout, prefix_grouper_enabled as _pg_enabled_fn, verify_on as _pg_verify_on, tol_ok as _pg_tol_ok, TOL_KILL as _PG_TOL_KILL\n"
+ " except Exception:\n"
+ " UNSLOTH_GRPO_PREFIX_GROUPER_ON = False\n"
+)
# Edit _get_per_token_logps to handle mixed precision
diff --git a/unsloth/utils/attention_dispatch.py b/unsloth/utils/attention_dispatch.py
index 2e984bad0a..68fb33dad9 100644
--- a/unsloth/utils/attention_dispatch.py
+++ b/unsloth/utils/attention_dispatch.py
@@ -17,6 +17,7 @@
from __future__ import annotations
+import os
from dataclasses import dataclass
from typing import Any, Optional, Tuple
@@ -42,6 +43,17 @@ if HAS_XFORMERS and torch.cuda.is_available():
HAS_XFORMERS = False
SDPA_HAS_GQA = "enable_gqa" in (scaled_dot_product_attention.__doc__ or "")
+# PrefixGrouper kernel, resolved once when the env gate is on so PG-off users never load
+# torch flex_attention.
+_flex_shared_prefix_attention = None
+if os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER", "1").lower() not in ("0", "false", "no", "off"):
+ try:
+ from .prefix_grouper_kernel import (
+ flex_shared_prefix_attention as _flex_shared_prefix_attention,
+ )
+ except Exception:
+ _flex_shared_prefix_attention = None
+
FLASH_VARLEN = "flash_varlen"
FLASH_DENSE = "flash_dense"
XFORMERS = "xformers"
@@ -84,6 +96,9 @@ class AttentionContext:
attention_mask: Optional[Tensor]
causal_mask: Optional[Any]
sliding_window: Optional[int] = None
+ # PrefixGrouper: non-None routes Q/K/V through the FlexAttention shared-prefix kernel;
+ # None leaves every existing construction/behavior unchanged.
+ prefix_seg_info: Optional[Any] = None
def select_attention_backend(use_varlen: bool = False) -> str:
@@ -99,6 +114,33 @@ def select_attention_backend(use_varlen: bool = False) -> str:
return SDPA
+def resolve_prefix_seg_info(kwargs, past_key_value, attention_mask):
+ """PrefixGrouper shared-prefix segment table resolver for the arch attention forwards.
+
+ The GRPO PrefixGrouper packed path rides a ``PrefixSegInfo`` in through ``**kwargs``
+ (same route as ``packed_seq_lengths``). When present, the forward must route Q/K/V
+ through the FlexAttention shared-prefix kernel via ``AttentionContext.prefix_seg_info``.
+
+ Returns the seg table (or ``None`` when PrefixGrouper did not group this batch -- the
+ unchanged path). Hardened: the shared-prefix stream is NOT a plain causal sequence, so running
+ it under a KV cache or an explicit padding mask would silently produce wrong logprobs.
+ That combination can only arise from misuse (PrefixGrouper only rides in via the GRPO
+ logprob forward, which is mask-free prefill), so we RAISE loudly instead of degrading
+ to a wrong result.
+
+ Factored here so every arch (llama/mistral/qwen3/gemma2/cohere/granite/falcon_h1)
+ shares one implementation and cannot drift.
+ """
+ seg = kwargs.get("prefix_seg_info", None)
+ if seg is not None and (past_key_value is not None or attention_mask is not None):
+ raise RuntimeError(
+ "PrefixGrouper: prefix_seg_info requires prefill with no KV cache and no "
+ f"attention_mask (got past_key_value={past_key_value is not None}, "
+ f"attention_mask={attention_mask is not None})."
+ )
+ return seg
+
+
def run_attention(
*, config: AttentionConfig, context: AttentionContext, Q: Tensor, K: Tensor, V: Tensor
) -> Tensor:
@@ -111,6 +153,28 @@ def run_attention(
and SDPA handle packing via a block-diagonal mask.
"""
+ # PrefixGrouper shared-prefix attention (GRPO dedup). Q/K/V here are [bsz, H, T, D];
+ # the kernel takes/returns [1, T, H, D], matching the other backends. The field is
+ # only set when the env gate is on and grouping succeeded; None keeps every backend
+ # byte-identical.
+ if context.prefix_seg_info is not None:
+ flex_shared_prefix_attention = _flex_shared_prefix_attention
+ if flex_shared_prefix_attention is None:
+ # gate flipped on after import (or one-time load failed): resolve lazily.
+ from ..utils.prefix_grouper_kernel import flex_shared_prefix_attention
+
+ scale = None
+ if config.flash_varlen_kwargs:
+ scale = config.flash_varlen_kwargs.get("softmax_scale")
+ A = flex_shared_prefix_attention(
+ Q.transpose(1, 2),
+ K.transpose(1, 2),
+ V.transpose(1, 2),
+ context.prefix_seg_info,
+ scale = scale,
+ )
+ return A # [1, T, n_heads, head_dim]
+
backend = config.backend
if backend == FLASH_VARLEN and context.seq_info is None:
backend = FLASH_DENSE if HAS_FLASH_ATTENTION else SDPA
@@ -337,5 +401,6 @@ __all__ = [
"AttentionConfig",
"AttentionContext",
"select_attention_backend",
+ "resolve_prefix_seg_info",
"run_attention",
]
diff --git a/unsloth/utils/prefix_grouper.py b/unsloth/utils/prefix_grouper.py
new file mode 100644
index 0000000000..4e6ff9672c
--- /dev/null
+++ b/unsloth/utils/prefix_grouper.py
@@ -0,0 +1,351 @@
+# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
+#
+# This program is free software: you can redistribute it and/or modify
+# it under the terms of the GNU Affero General Public License as published by
+# the Free Software Foundation, either version 3 of the License, or
+# (at your option) any later version.
+#
+# This program is distributed in the hope that it will be useful,
+# but WITHOUT ANY WARRANTY; without even the implied warranty of
+# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
+# GNU Affero General Public License for more details.
+#
+# You should have received a copy of the GNU Affero General Public License
+# along with this program. If not, see .
+
+"""PrefixGrouper layout builder + completion-logprob extraction for the Unsloth GRPO
+packed path (all archs that route through the varlen attention dispatch).
+
+Given the de-padded, LEFT-PACKED input_ids the packed GRPO path already works with, this
+module:
+
+ 1. Detects consecutive ``num_generations`` rows that share a prompt prefix (byte-
+ identical prompt precondition; falls back / returns None otherwise).
+ 2. Builds ONE flat shared-prefix stream across all groups
+ ``[ prefix_g0, suf_g0_0 .. suf_g0_{G-1}, prefix_g1, ... ]`` with position_ids that
+ continue each prefix positionally, plus a ``PrefixSegInfo`` segment table for the
+ FlexAttention shared-prefix kernel.
+ 3. Extracts completion logprobs via the index map (completion pos ``j==0`` predicted
+ from the shared prefix's last token; ``j>=1`` from the preceding suffix token) and
+ scatters them back into ``[total_rows, W]`` EXACTLY where the full-row packed path
+ puts them (dest = ``orig_row*L + orig_col``), so grpo_compute_loss / completion_mask
+ / TIS / metrics are byte-untouched.
+
+The flat stream is built by GATHERING original (row, col) coordinates out of input_ids,
+so the grad path's autograd flows to the same embedding rows as today (the shared prefix
+now contributes grad once = the sum of the G repeats, which is mathematically identical).
+
+``chunked_hidden_states_selective_log_softmax`` (from unsloth_zoo, passed in) is reused
+verbatim over the gathered predicting-position hidden states, so fp32 accumulation,
+logit_scale/softcapping/temperature are all preserved.
+
+Env:
+ UNSLOTH_GRPO_PREFIX_GROUPER=1 engage (default ON; set 0 to disable). Auto-off under vLLM.
+ UNSLOTH_GRPO_PREFIX_GROUPER_TOKR=1.3 tok_r auto-gate threshold (env-overridable)
+ UNSLOTH_GRPO_PREFIX_GROUPER_VERIFY=1 first-step self-verify (default ON)
+ UNSLOTH_GRPO_PREFIX_GROUPER_TOL=0.7 self-verify PASS band (nats)
+"""
+
+from __future__ import annotations
+
+import os
+from dataclasses import dataclass
+from typing import List, Optional, Tuple
+
+import torch
+
+from .prefix_grouper_kernel import build_seg_info_multigroup, PrefixSegInfo
+
+
+# ---------------------------------------------------------------------------
+# Env helpers
+# ---------------------------------------------------------------------------
+def env_on(name: str, default: str = "0") -> bool:
+ return os.environ.get(name, default).lower() not in ("0", "false", "no", "off")
+
+
+# One-time env reads; the helpers stay callable since unsloth_zoo imports and calls them.
+_ENABLED = env_on("UNSLOTH_GRPO_SEQ_PACKING", "1") and env_on("UNSLOTH_GRPO_PREFIX_GROUPER", "1")
+_VERIFY_ON = env_on("UNSLOTH_GRPO_PREFIX_GROUPER_VERIFY", "1")
+_TOKR_THRESHOLD = float(os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER_TOKR", "1.3"))
+_TOL_OK = float(os.environ.get("UNSLOTH_GRPO_PREFIX_GROUPER_TOL", "0.7"))
+
+
+def prefix_grouper_enabled() -> bool:
+ """PrefixGrouper requires seq-packing on (it reuses its de-pad + scatter machinery)."""
+ return _ENABLED
+
+
+def verify_on() -> bool:
+ return _VERIFY_ON
+
+
+def tokr_threshold() -> float:
+ return _TOKR_THRESHOLD
+
+
+def tol_ok() -> float:
+ return _TOL_OK
+
+
+# diff >= TOL_KILL = broken mask/isolation -> structure permanently unsafe; between
+# tol_ok and TOL_KILL -> fall back for this shape but keep trying others.
+TOL_KILL = 1.5
+
+
+@dataclass
+class GroupLayout:
+ """Everything the GRPO forward needs to run + extract the shared-prefix path."""
+
+ flat_ids: torch.Tensor # [1, T] (T == seg.T)
+ position_ids: torch.Tensor # [1, T]
+ prefix_seg_info: PrefixSegInfo
+ # per completion target token, aligned 1:1:
+ tgt_rows: torch.Tensor # [N] original row index
+ tgt_cols: torch.Tensor # [N] original padded column in that row
+ tgt_pred: torch.Tensor # [N] flat predicting index (into the T stream)
+ tgt_flat: torch.Tensor # [N] flat index of the target token itself (into T)
+ total_rows: int
+ L: int # original padded seq length (input_ids.shape[1])
+ W: int # logits_to_keep + max_left_pad (scatter width)
+ tok_r: float
+ signature: Tuple
+
+ def extract_logps(
+ self,
+ hidden,
+ lm_head,
+ chunked_fn,
+ chunks,
+ logit_scale_multiply,
+ logit_scale_divide,
+ logit_softcapping,
+ temperature,
+ ) -> torch.Tensor:
+ """hidden: [1, T, Hdim] (pre-lm_head hidden states, UNSLOTH_RETURN_HIDDEN_STATES=1).
+ Returns [total_rows, W] float32, byte-compatible with the packed path result."""
+ # In a sharded model hidden may live on the lm-head device; move the small index
+ # maps to hidden.device before indexing.
+ device = hidden.device
+ pred_h = hidden[0, self.tgt_pred.to(device), :].unsqueeze(0) # [1, N, Hdim]
+ tgt_ids = self.flat_ids[0, self.tgt_flat].to(device).unsqueeze(0) # [1, N]
+ sel = chunked_fn(
+ pred_h,
+ lm_head,
+ tgt_ids,
+ chunks,
+ logit_scale_multiply,
+ logit_scale_divide,
+ logit_softcapping,
+ temperature,
+ )[0] # [N] logprobs
+ dest = self.tgt_rows.to(device) * self.L + self.tgt_cols.to(device)
+ result = (
+ torch.zeros(self.total_rows * self.L, dtype = torch.float32, device = device)
+ .index_put((dest,), sel.to(torch.float32))
+ .view(self.total_rows, self.L)[:, -self.W :]
+ )
+ return result
+
+
+def _build_groups(ids_cpu, real_cols_cpu, cstart_cpu, num_generations, total_rows):
+ """CPU-side grouping. Returns group dicts or None. Mirrors the packed _pk_* partition.
+
+ A row's REAL tokens are the columns where input != pad. Its completion region (what
+ the packed path scatters, then completion_mask masks) is the real columns with
+ original col >= cstart_r, where cstart_r = (L - logits_to_keep) - left_pad_r. The
+ prompt is the real columns < cstart_r. Within a GRPO group all G rows share the same
+ prompt => same left_pad => same cstart => the prompt real columns are BYTE-IDENTICAL
+ across the group (the shared prefix). We require that byte-identity (falls back
+ otherwise). No prompt-tail special-casing: every suffix token is scattered exactly
+ like the packed path; completion_mask masks the leading prompt-tail positions.
+ """
+ G = num_generations
+ if G is None or G < 2 or total_rows % G != 0:
+ return None
+ groups = []
+ for g0 in range(0, total_rows, G):
+ rows = list(range(g0, g0 + G))
+ prompt_cols_per_row = [] # real cols < cstart
+ prompt_toks_per_row = []
+ comp_cols_per_row = [] # real cols >= cstart (the completion region packed scatters)
+ for r in rows:
+ cs = cstart_cpu[r]
+ rc = real_cols_cpu[r]
+ p_cols = [c for c in rc if c < cs]
+ c_cols = [c for c in rc if c >= cs]
+ prompt_cols_per_row.append(p_cols)
+ prompt_toks_per_row.append([ids_cpu[r][c] for c in p_cols])
+ comp_cols_per_row.append(c_cols)
+ if any(len(p) == 0 for p in prompt_toks_per_row):
+ return None
+ # require BYTE-IDENTICAL prompts across the group (shared-prefix precondition).
+ P = len(prompt_toks_per_row[0])
+ if any(len(prompt_toks_per_row[k]) != P for k in range(1, G)):
+ return None
+ p0 = prompt_toks_per_row[0]
+ if any(prompt_toks_per_row[k] != p0 for k in range(1, G)):
+ return None
+ if P == 0:
+ return None
+ R_list = [len(c) for c in comp_cols_per_row]
+ if sum(R_list) == 0:
+ return None
+ groups.append(
+ dict(
+ rows = rows,
+ P = P,
+ prefix_cols = prompt_cols_per_row[0], # shared prompt real columns (row0)
+ prefix_row = rows[0],
+ R_list = R_list,
+ suf_cols = comp_cols_per_row, # per-row completion-region real columns
+ )
+ )
+ return groups
+
+
+def _tok_r(groups) -> float:
+ tok_full = 0
+ tok_sp = 0
+ for gm in groups:
+ P = gm["P"]
+ Rs = gm["R_list"]
+ tok_full += sum(P + r for r in Rs) # G*P + sumR
+ tok_sp += P + sum(Rs) # P + sumR
+ return (tok_full / tok_sp) if tok_sp else 1.0
+
+
+def build_group_layout(
+ input_ids,
+ logits_to_keep,
+ pad_id,
+ num_generations,
+ left_pad_tokens_per_prompt,
+ *,
+ apply_tokr_gate = True,
+ max_segment_cap = None,
+):
+ """Build the shared-prefix GroupLayout, or return None to fall back to the packed path.
+
+ input_ids : [B, L]. GRPO's layout is left-padded in the prompt and right-padded in
+ the completion. Real tokens of a row are a contiguous run not necessarily
+ starting at column 0.
+ logits_to_keep : int
+ left_pad_tokens_per_prompt : [B] long tensor (per-row left-pad count in the prompt).
+ """
+ device = input_ids.device
+ total_rows, L = input_ids.shape
+ keep = input_ids != pad_id
+ # completion start column per row (matches create_completion_attention_mask / _pk_cstart).
+ cstart = ((L - logits_to_keep) - left_pad_tokens_per_prompt).to(torch.long)
+ cstart_cpu = cstart.tolist()
+ ids_cpu = input_ids.tolist()
+ # per-row real (non-pad) columns. GRPO rows are one contiguous real run, so derive
+ # [first, first+n) on GPU; the O(B*L) scan is only a non-contiguous fallback.
+ n_real = keep.sum(dim = 1)
+ first = torch.argmax(keep.to(torch.int8), dim = 1)
+ ar = torch.arange(L, device = device)
+ contiguous = bool(
+ (keep == ((ar >= first.unsqueeze(1)) & (ar < (first + n_real).unsqueeze(1)))).all()
+ )
+ if contiguous:
+ real_cols_cpu = [list(range(f, f + n)) for f, n in zip(first.tolist(), n_real.tolist())]
+ else:
+ keep_cpu = keep.tolist()
+ real_cols_cpu = [[c for c in range(L) if keep_cpu[r][c]] for r in range(total_rows)]
+
+ groups = _build_groups(ids_cpu, real_cols_cpu, cstart_cpu, num_generations, total_rows)
+ if groups is None:
+ return None
+
+ # sliding-window guard: a group's PG span is P + max(R); fall back if it exceeds the window.
+ if max_segment_cap is not None:
+ for gm in groups:
+ if gm["P"] + max(gm["R_list"]) > max_segment_cap:
+ return None
+
+ tok_r = _tok_r(groups)
+ if apply_tokr_gate and tok_r < tokr_threshold():
+ return None # low reuse -> not worth it; use the full-row packed path
+
+ # Build flat stream by gathering original (row, col) coordinates.
+ group_specs = [(gm["P"], gm["R_list"]) for gm in groups]
+ seg, group_meta = build_seg_info_multigroup(group_specs, device)
+
+ flat_src_rows: List[int] = []
+ flat_src_cols: List[int] = []
+ pos_list: List[int] = []
+ tgt_rows: List[int] = []
+ tgt_cols: List[int] = []
+ tgt_pred: List[int] = []
+ tgt_flat: List[int] = []
+
+ for gm, meta in zip(groups, group_meta):
+ rows = gm["rows"]
+ P = gm["P"]
+ r0 = gm["prefix_row"]
+ prefix_cols = gm["prefix_cols"] # ORIGINAL real prompt columns (len P) of row0
+ plast = meta["prefix_last_index"] # base + P - 1
+ # gather the shared prefix once, from row0.
+ flat_src_rows.extend([r0] * P)
+ flat_src_cols.extend(prefix_cols)
+ pos_list.extend(range(P))
+ # suffixes: every suffix token is a completion-region target (scattered like the
+ # packed path; completion_mask hides prompt-tail positions).
+ for i, r in enumerate(rows):
+ cols = gm["suf_cols"][i]
+ r_i = len(cols)
+ s, e = meta["suffix_slices"][i] # flat offsets [s, e)
+ flat_src_rows.extend([r] * r_i)
+ flat_src_cols.extend(cols)
+ pos_list.extend(range(P, P + r_i))
+ for j in range(r_i):
+ # pos 0 is predicted from the prefix's last token; j>=1 from the previous suffix token.
+ pred = plast if j == 0 else (s + j - 1)
+ tgt_rows.append(r)
+ tgt_cols.append(cols[j]) # ORIGINAL padded column in row r
+ tgt_pred.append(pred)
+ tgt_flat.append(s + j) # flat index of the target token itself
+
+ T = len(flat_src_rows)
+ assert T == seg.T, f"flat stream len {T} != seg.T {seg.T}"
+ fr = torch.tensor(flat_src_rows, device = device, dtype = torch.long)
+ fc = torch.tensor(flat_src_cols, device = device, dtype = torch.long)
+ flat_ids = input_ids[fr, fc].unsqueeze(0) # [1, T] (grad-safe gather)
+ position_ids = torch.tensor(pos_list, device = device, dtype = torch.long).unsqueeze(0)
+
+ max_left_pad = int(left_pad_tokens_per_prompt.max().item()) if total_rows else 0
+ W = logits_to_keep + max_left_pad
+
+ # self-verify cache key: the mask/index-map/scatter logic is structural, so key on
+ # (num_groups, group_sizes), not exact lengths -- GRPO lengths change every step and
+ # keying on T would re-verify forever ("verify once, then trust", like the packed path).
+ grp_sizes = tuple(sorted(len(gm["R_list"]) for gm in groups))
+ sig = (len(groups), grp_sizes)
+
+ return GroupLayout(
+ flat_ids = flat_ids,
+ position_ids = position_ids,
+ prefix_seg_info = seg,
+ tgt_rows = torch.tensor(tgt_rows, device = device, dtype = torch.long),
+ tgt_cols = torch.tensor(tgt_cols, device = device, dtype = torch.long),
+ tgt_pred = torch.tensor(tgt_pred, device = device, dtype = torch.long),
+ tgt_flat = torch.tensor(tgt_flat, device = device, dtype = torch.long),
+ total_rows = total_rows,
+ L = L,
+ W = W,
+ tok_r = tok_r,
+ signature = sig,
+ )
+
+
+__all__ = [
+ "GroupLayout",
+ "build_group_layout",
+ "prefix_grouper_enabled",
+ "verify_on",
+ "tokr_threshold",
+ "tol_ok",
+ "TOL_KILL",
+ "env_on",
+]
diff --git a/unsloth/utils/prefix_grouper_kernel.py b/unsloth/utils/prefix_grouper_kernel.py
new file mode 100644
index 0000000000..9a9719b015
--- /dev/null
+++ b/unsloth/utils/prefix_grouper_kernel.py
@@ -0,0 +1,436 @@
+# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
+#
+# This program is free software: you can redistribute it and/or modify
+# it under the terms of the GNU Affero General Public License as published by
+# the Free Software Foundation, either version 3 of the License, or
+# (at your option) any later version.
+#
+# This program is distributed in the hope that it will be useful,
+# but WITHOUT ANY WARRANTY; without even the implied warranty of
+# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
+# GNU Affero General Public License for more details.
+#
+# You should have received a copy of the GNU Affero General Public License
+# along with this program. If not, see .
+
+"""FlexAttention shared-prefix kernel for PrefixGrouper (GRPO shared-prompt dedup).
+
+In GRPO every prompt spawns ``G = num_generations`` completions that share the same
+prompt prefix. The full-row packed path forwards the identical prefix ``G`` times.
+PrefixGrouper stores the prefix ONCE and concatenates only the ``G`` suffixes, with an
+attention layout where each suffix token attends to ``[the single shared prefix] +
+[causal within its own suffix]``. This kernel expresses that one-prefix -> many-suffix
+fan-out via a ``torch.nn.attention.flex_attention`` block mask, so the masked-out
+cross-suffix / cross-group blocks are never computed and the ``P + G*R`` FLOP saving is
+realised (not merely a masked dense ``O(T^2)``).
+
+Mask semantics (identical to the certified SDPA oracle):
+
+ keep(q_idx, kv_idx) = same_group(q, kv) AND
+ ( is_prefix[kv_idx] # full prefix visibility
+ OR ( suffix_of_kv[kv_idx] == suffix_of_kv[q_idx] # same suffix ...
+ AND kv_idx <= q_idx ) ) # ... causal within it
+
+This module is self-contained (no dependency on any temp/ scratch dir) so PrefixGrouper
+works from the installed source after a fresh compile. It is only imported lazily from
+``attention_dispatch.run_attention`` when ``prefix_seg_info`` is present, which itself is
+only ever set when ``UNSLOTH_GRPO_PREFIX_GROUPER`` is on and grouping succeeded, so the
+default (off) path never touches this file.
+
+Provided entry points:
+ * ``PrefixSegInfo`` : per-flat-token segment metadata + cache signature.
+ * ``build_seg_info_multigroup``: build PrefixSegInfo for many groups packed flat.
+ * ``build_seg_info_from_layout``: build PrefixSegInfo for ONE group (test helper).
+ * ``get_block_mask`` : cached create_block_mask keyed on the signature.
+ * ``flex_shared_prefix_attention(Q, K, V, prefix_seg_info)``
+ Q/K/V of shape [1, T, n_heads, head_dim]; returns [1, T, n_heads, head_dim],
+ IDENTICAL semantics to the SDPA oracle.
+"""
+
+from __future__ import annotations
+
+import os
+from dataclasses import dataclass
+from typing import Dict, List, Optional, Tuple
+
+import torch
+from torch.nn.attention.flex_attention import (
+ BlockMask,
+ create_block_mask,
+ flex_attention,
+)
+
+# GRPO feeds many distinct segment lengths; at dynamo's default recompile_limit (8) the
+# compiled kernel silently reuses a mismatched specialisation (wrong results). Raise it.
+torch._dynamo.config.recompile_limit = max(getattr(torch._dynamo.config, "recompile_limit", 8), 256)
+torch._dynamo.config.accumulated_recompile_limit = max(
+ getattr(torch._dynamo.config, "accumulated_recompile_limit", 256), 2048
+)
+
+
+# Compiled kernels: torch.compile fuses the sparse mask into one kernel. dynamic=True is
+# required: T changes almost every GRPO batch and dynamic=False recompiles per T (~14s
+# each). T is still padded to a multiple of 128 (_pad_len) for the backward kernel.
+_flex_attention_compiled = torch.compile(flex_attention, dynamic = True)
+_create_block_mask_compiled = torch.compile(create_block_mask, dynamic = True)
+
+# Flash block sizes by Q dtype (env-overridable). The two disjoint key runs (prefix +
+# own-suffix) stress online-softmax accumulation: fp32 needs 32/32 for a ~1e-6 floor;
+# bf16 passes parity at 128/64 and is ~5x faster (128/128 OOMs Triton on B200).
+_FP32_BLOCK_M = int(os.environ.get("PG_FLEX_BLOCK_M", "32"))
+_FP32_BLOCK_N = int(os.environ.get("PG_FLEX_BLOCK_N", "32"))
+_BF16_BLOCK_M = int(os.environ.get("PG_FLEX_BF16_BLOCK_M", "128"))
+_BF16_BLOCK_N = int(os.environ.get("PG_FLEX_BF16_BLOCK_N", "64"))
+
+
+def _kernel_options_for_dtype(dtype):
+ """Pick the numerically-safe flash block sizes for the Q dtype."""
+ if dtype == torch.bfloat16 or dtype == torch.float16:
+ return {"BLOCK_M": _BF16_BLOCK_M, "BLOCK_N": _BF16_BLOCK_N}
+ return {"BLOCK_M": _FP32_BLOCK_M, "BLOCK_N": _FP32_BLOCK_N}
+
+
+# Backward-compat constant (fp32 default).
+_FLEX_KERNEL_OPTIONS = {"BLOCK_M": _FP32_BLOCK_M, "BLOCK_N": _FP32_BLOCK_N}
+
+# The compiled backward trips an Inductor assertion when T is not a multiple of 128, so
+# pad the flat sequence. Pad tokens form a group that attends to / is attended by nothing
+# (all-masked rows return 0, not NaN) and are sliced off the output.
+_PAD_MULTIPLE = 128
+_PAD_GROUP = -99 # sentinel group id / suffix id for pad tokens
+
+
+def _pad_len(T: int) -> int:
+ return ((T + _PAD_MULTIPLE - 1) // _PAD_MULTIPLE) * _PAD_MULTIPLE
+
+
+# ---------------------------------------------------------------------------
+# Segment metadata
+# ---------------------------------------------------------------------------
+
+
+@dataclass
+class PrefixSegInfo:
+ """Per-flat-token segment metadata driving the shared-prefix block mask.
+
+ The label tensors are 1-D of length ``T_pad`` (>= real ``T``, padded up to a multiple
+ of 128 so the backward kernel compiles). Positions ``[T:T_pad)`` are pad tokens
+ (group/suffix == _PAD_GROUP) that attend to nothing.
+
+ Attributes
+ ----------
+ group_of_kv : LongTensor [T_pad]
+ Group id per flat token (0..num_groups-1); _PAD_GROUP for pad tokens.
+ is_prefix : BoolTensor [T_pad]
+ True iff the token is a prefix token of its group (False for pad).
+ suffix_of_kv : LongTensor [T_pad]
+ Suffix id per flat token; -1 for prefix, _PAD_GROUP for pad. Suffix ids are
+ globally unique across groups.
+ signature : hashable
+ Cache key for the block mask (depends only on the labels + T_pad).
+ T : int
+ Real flat sequence length (Q/K/V of this length are padded internally).
+ T_pad : int
+ Padded length (multiple of 128) at which the block mask is built.
+ """
+
+ group_of_kv: torch.Tensor
+ is_prefix: torch.Tensor
+ suffix_of_kv: torch.Tensor
+ signature: Tuple
+ T: int
+ T_pad: int
+
+
+def _pad_labels(group_of_kv, is_prefix, suffix_of_kv, device):
+ """Pad the label tensors up to a multiple of 128 with pad-token sentinels."""
+ T = int(group_of_kv.numel())
+ T_pad = _pad_len(T)
+ if T_pad == T:
+ return group_of_kv, is_prefix, suffix_of_kv, T, T_pad
+ pad = T_pad - T
+ group_of_kv = torch.cat(
+ [group_of_kv, torch.full((pad,), _PAD_GROUP, dtype = torch.long, device = device)]
+ )
+ is_prefix = torch.cat([is_prefix, torch.zeros(pad, dtype = torch.bool, device = device)])
+ suffix_of_kv = torch.cat(
+ [suffix_of_kv, torch.full((pad,), _PAD_GROUP, dtype = torch.long, device = device)]
+ )
+ return group_of_kv, is_prefix, suffix_of_kv, T, T_pad
+
+
+def build_seg_info_from_layout(layout, device: Optional[torch.device] = None) -> PrefixSegInfo:
+ """Build PrefixSegInfo for ONE group from an object with ``.flat_ids``, ``.P`` and
+ ``.suffix_slices`` (used by the parity test / oracle helpers)."""
+ if device is None:
+ device = layout.flat_ids.device
+ T = int(layout.flat_ids.shape[1])
+ P = int(layout.P)
+
+ group_of_kv = torch.zeros(T, dtype = torch.long, device = device) # single group -> 0
+ is_prefix = torch.zeros(T, dtype = torch.bool, device = device)
+ is_prefix[:P] = True
+ suffix_of_kv = torch.full((T,), -1, dtype = torch.long, device = device)
+ for i, (s, e) in enumerate(layout.suffix_slices):
+ suffix_of_kv[s:e] = i
+
+ group_of_kv, is_prefix, suffix_of_kv, T, T_pad = _pad_labels(
+ group_of_kv, is_prefix, suffix_of_kv, device
+ )
+ sig = ("single", T_pad, P, tuple((s, e) for (s, e) in layout.suffix_slices))
+ return PrefixSegInfo(
+ group_of_kv = group_of_kv,
+ is_prefix = is_prefix,
+ suffix_of_kv = suffix_of_kv,
+ signature = sig,
+ T = T,
+ T_pad = T_pad,
+ )
+
+
+def build_seg_info_multigroup(
+ group_specs: List[Tuple[int, List[int]]], device: torch.device
+) -> Tuple[PrefixSegInfo, List[dict]]:
+ """Build PrefixSegInfo for several shared-prefix groups packed block-diagonally.
+
+ Parameters
+ ----------
+ group_specs : list of (P_g, [R_{g,0}, R_{g,1}, ...])
+ For each group: prefix length and the list of suffix lengths.
+
+ Returns
+ -------
+ seg : PrefixSegInfo
+ group_meta : list of dicts with 'base', 'P', 'prefix_last_index', 'suffix_slices'
+ (flat offsets), enough to build the completion index map.
+ """
+ group_of_list = []
+ is_prefix_list = []
+ suffix_of_list = []
+ group_meta = []
+
+ base = 0
+ suffix_counter = 0
+ sig_parts = []
+ for gid, (P, R_list) in enumerate(group_specs):
+ # prefix
+ group_of_list.append(torch.full((P,), gid, dtype = torch.long, device = device))
+ is_prefix_list.append(torch.ones(P, dtype = torch.bool, device = device))
+ suffix_of_list.append(torch.full((P,), -1, dtype = torch.long, device = device))
+ prefix_last_index = base + P - 1
+ suffix_slices = []
+ cursor = base + P
+ for r in R_list:
+ group_of_list.append(torch.full((r,), gid, dtype = torch.long, device = device))
+ is_prefix_list.append(torch.zeros(r, dtype = torch.bool, device = device))
+ suffix_of_list.append(torch.full((r,), suffix_counter, dtype = torch.long, device = device))
+ suffix_slices.append((cursor, cursor + r))
+ cursor += r
+ suffix_counter += 1
+ group_meta.append(
+ {
+ "base": base,
+ "P": P,
+ "prefix_last_index": prefix_last_index,
+ "suffix_slices": suffix_slices,
+ }
+ )
+ sig_parts.append((P, tuple(R_list)))
+ base = cursor
+
+ group_of_kv = torch.cat(group_of_list)
+ is_prefix = torch.cat(is_prefix_list)
+ suffix_of_kv = torch.cat(suffix_of_list)
+ group_of_kv, is_prefix, suffix_of_kv, T, T_pad = _pad_labels(
+ group_of_kv, is_prefix, suffix_of_kv, device
+ )
+ sig = ("multi", T_pad, tuple(sig_parts))
+ seg = PrefixSegInfo(
+ group_of_kv = group_of_kv,
+ is_prefix = is_prefix,
+ suffix_of_kv = suffix_of_kv,
+ signature = sig,
+ T = T,
+ T_pad = T_pad,
+ )
+ return seg, group_meta
+
+
+# ---------------------------------------------------------------------------
+# Block-mask builder + cache, keyed on (signature, device): the mask depends only on the
+# per-token labels and T, so it is reused across layers and steps.
+
+_BLOCK_MASK_CACHE: Dict[Tuple, BlockMask] = {}
+
+
+def _make_mask_mod(group_of_kv, is_prefix, suffix_of_kv):
+ """Return a mask_mod closure over the (device) label tensors.
+
+ keep(q, kv) = same_group AND
+ ( is_prefix[kv] AND kv <= q # causal within/ into prefix
+ OR ( suffix_of_kv[kv] == suffix_of_kv[q] # same suffix ...
+ AND (not is_prefix[q]) # q is a suffix token ...
+ AND kv <= q ) ) # ... causal within it
+
+ The single ``kv <= q`` guard on the is_prefix branch gives BOTH prefix-causal
+ behaviour (a prefix q sees only earlier prefix tokens) AND full-prefix-visibility for
+ suffixes (every prefix index < every suffix index in a group, so kv <= q always holds
+ for a suffix q vs a prefix kv of its group), matching the SDPA oracle exactly.
+ """
+
+ def mask_mod(b, h, q_idx, kv_idx):
+ same_group = group_of_kv[q_idx] == group_of_kv[kv_idx]
+ kv_is_prefix = is_prefix[kv_idx]
+ causal = kv_idx <= q_idx
+ same_suffix = (suffix_of_kv[kv_idx] == suffix_of_kv[q_idx]) & (~is_prefix[q_idx])
+ keep = same_group & ((kv_is_prefix & causal) | (same_suffix & causal))
+ return keep
+
+ return mask_mod
+
+
+def get_block_mask(
+ seg: PrefixSegInfo,
+ device: torch.device,
+ compile_mask: bool = True,
+) -> BlockMask:
+ """Return a cached BlockMask for the segment signature (built once, reused).
+
+ CRITICAL: the block mask is cached and shared across BOTH the no-grad old/ref logprob
+ forward (which runs under torch.inference_mode) and the grad training forward. If the
+ mask were first built under inference_mode, its tensors would be INFERENCE tensors that
+ "cannot be saved for backward" when reused in the grad forward. We therefore build the
+ mask with inference mode explicitly DISABLED, so the same cached BlockMask is a normal
+ tensor usable by autograd. (The mask depends only on integer labels; it needs no grad.)
+ """
+ key = (seg.signature, str(device))
+ bm = _BLOCK_MASK_CACHE.get(key)
+ if bm is not None:
+ return bm
+
+ # Move labels to the consumer (Q) device: with a sharded model the seg tensors live on
+ # input_ids.device and would index cross-device. Copies once per (signature, device).
+ # These copies must also run with inference mode DISABLED (same reason as the mask build):
+ # when this entry is first built under the no-grad old/ref forward's inference_mode and
+ # device != seg.device, a .to(device) copy would be an inference tensor that mask_mod
+ # captures, which then cannot be saved for backward when the grad training forward reuses
+ # the cached mask.
+ builder = _create_block_mask_compiled if compile_mask else create_block_mask
+ with torch.inference_mode(False):
+ mask_mod = _make_mask_mod(
+ seg.group_of_kv.to(device), seg.is_prefix.to(device), seg.suffix_of_kv.to(device)
+ )
+ bm = builder(
+ mask_mod,
+ B = 1,
+ H = None,
+ Q_LEN = seg.T_pad,
+ KV_LEN = seg.T_pad,
+ device = device,
+ )
+ # FIFO bound: GRPO lengths change nearly every step, so evict the oldest to cap GPU pins.
+ if len(_BLOCK_MASK_CACHE) >= 8:
+ _BLOCK_MASK_CACHE.pop(next(iter(_BLOCK_MASK_CACHE)))
+ _BLOCK_MASK_CACHE[key] = bm
+ return bm
+
+
+def clear_block_mask_cache():
+ _BLOCK_MASK_CACHE.clear()
+
+
+def _pad_qkv_seq(x: torch.Tensor, T_pad: int) -> torch.Tensor:
+ """Zero-pad a [B, H, T, D] tensor along the sequence dim up to T_pad."""
+ T = x.shape[2]
+ if T_pad == T:
+ return x
+ pad = torch.zeros(x.shape[0], x.shape[1], T_pad - T, x.shape[3], device = x.device, dtype = x.dtype)
+ return torch.cat([x, pad], dim = 2)
+
+
+def _run_flex(q, k, v, block_mask, enable_gqa, scale, compiled, T, T_pad):
+ """Pad q/k/v to T_pad, run flex, slice the output back to T. q/k/v: [B,H,T,D]."""
+ qp = _pad_qkv_seq(q, T_pad)
+ kp = _pad_qkv_seq(k, T_pad)
+ vp = _pad_qkv_seq(v, T_pad)
+ if compiled:
+ out = _flex_attention_compiled(
+ qp,
+ kp,
+ vp,
+ block_mask = block_mask,
+ enable_gqa = enable_gqa,
+ scale = scale,
+ kernel_options = _kernel_options_for_dtype(qp.dtype),
+ )
+ else:
+ # eager path (fp64 parity): dense scores, no kernel_options.
+ out = flex_attention(
+ qp,
+ kp,
+ vp,
+ block_mask = block_mask,
+ enable_gqa = enable_gqa,
+ scale = scale,
+ )
+ return out[:, :, :T, :]
+
+
+# ---------------------------------------------------------------------------
+# The kernel entry point
+# ---------------------------------------------------------------------------
+
+
+def flex_shared_prefix_attention(
+ Q: torch.Tensor,
+ K: torch.Tensor,
+ V: torch.Tensor,
+ prefix_seg_info: PrefixSegInfo,
+ scale: Optional[float] = None,
+ block_mask: Optional[BlockMask] = None,
+ compiled: bool = True,
+) -> torch.Tensor:
+ """Shared-prefix attention via FlexAttention.
+
+ Parameters
+ ----------
+ Q, K, V : Tensor [1, T, n_heads, head_dim]
+ (Q has n_heads, K/V have n_kv_heads for GQA).
+ prefix_seg_info : PrefixSegInfo
+ scale : optional float, softmax scale (defaults to 1/sqrt(head_dim)).
+ block_mask : optional precomputed BlockMask (else built/cached from seg info).
+
+ Returns
+ -------
+ Tensor [1, T, n_heads, head_dim], identical semantics to the SDPA oracle branch.
+ """
+ assert Q.dim() == 4 and Q.shape[0] == 1, f"expected [1,T,H,D], got {tuple(Q.shape)}"
+ device = Q.device
+ # FlexAttention wants [B, H, T, D].
+ q = Q.transpose(1, 2) # [1, n_heads, T, D]
+ k = K.transpose(1, 2) # [1, n_kv_heads, T, D]
+ v = V.transpose(1, 2)
+
+ n_heads = q.shape[1]
+ n_kv = k.shape[1]
+ enable_gqa = n_heads != n_kv
+ T = q.shape[2]
+ T_pad = prefix_seg_info.T_pad
+ assert T == prefix_seg_info.T, f"Q length {T} != seg.T {prefix_seg_info.T}"
+
+ if block_mask is None:
+ block_mask = get_block_mask(prefix_seg_info, device, compile_mask = compiled)
+
+ out = _run_flex(q, k, v, block_mask, enable_gqa, scale, compiled, T, T_pad)
+ # back to [1, T, n_heads, D]
+ return out.transpose(1, 2).contiguous()
+
+
+__all__ = [
+ "PrefixSegInfo",
+ "build_seg_info_multigroup",
+ "build_seg_info_from_layout",
+ "get_block_mask",
+ "clear_block_mask_cache",
+ "flex_shared_prefix_attention",
+]