File size: 21,938 Bytes
b22e03e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
285871e
af9e65b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
285871e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b22e03e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af9e65b
285871e
b22e03e
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
"""FlashRT Flex-style block-sparse attention training API.

The public function implements the PI052 prefix/action mask pattern:

* prefix query rows use the original K/V tensors, so prefix losses keep normal
  gradients into prefix K/V;
* action query rows read detached prefix K/V plus normal action K/V by default,
  matching the current training semantics.

Unsupported shapes route to the SDPA reference path. Native CUDA kernels are
not exposed until a shape-specialized implementation beats SDPA on the target
A100/5090 validation matrix.
"""

from __future__ import annotations

from typing import Optional

import torch
import torch.nn.functional as F

try:
    from ._ops import ops

    _HAS_OPS = hasattr(ops, "_flashrt_training_package_marker")
except Exception:  # source-tree tests before kernel-builder creates _ops.py
    ops = None
    _HAS_OPS = False


MASK_VALUE_F32 = -2.3819763e38


def _use_ops(namespace_ops) -> None:
    """Install a manually built extension (dev/testing path)."""
    global ops, _HAS_OPS
    ops = namespace_ops
    _HAS_OPS = hasattr(ops, "_flashrt_training_package_marker")


def _check_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> None:
    if q.dim() != 4 or k.dim() != 4 or v.dim() != 4:
        raise ValueError("q, k, and v must be shaped (B, H, S, D)")
    if q.shape[0] != k.shape[0] or q.shape[0] != v.shape[0]:
        raise ValueError("q, k, and v batch dimensions must match")
    if k.shape != v.shape:
        raise ValueError("k and v shapes must match")
    if q.shape[2] != k.shape[2] or q.shape[3] != k.shape[3]:
        raise ValueError("q, k, and v sequence/head_dim dimensions must match")
    if q.device != k.device or q.device != v.device:
        raise ValueError("q, k, and v must be on the same device")


def _as_valid(mask: Optional[torch.Tensor], batch: int, length: int, device: torch.device) -> torch.Tensor:
    if mask is None:
        return torch.ones((batch, length), dtype=torch.bool, device=device)
    if mask.shape != (batch, length):
        raise ValueError(f"mask must be shaped {(batch, length)}, got {tuple(mask.shape)}")
    return mask.to(device=device, dtype=torch.bool)


def build_block_sparse_bool_masks(
    prefix_valid: Optional[torch.Tensor],
    prefix_att: Optional[torch.Tensor],
    *,
    batch: int,
    prefix_len: int,
    action_len: int,
    action_block_size: int,
    non_fast_prefix_len: Optional[int] = None,
    action_valid: Optional[torch.Tensor] = None,
    device: Optional[torch.device] = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Build boolean masks for the split FlexAttention SDPA calls.

    Returns ``(prefix_rows, action_rows)`` with shapes ``(B, P, S)`` and
    ``(B, A, S)``. Boolean True means the key/value position is visible.

    ``prefix_att`` follows Lerobot's cumulative-block convention: prefix key
    ``j`` is visible to prefix query ``i`` when ``cumsum(prefix_att)[j] <=
    cumsum(prefix_att)[i]`` and both rows are valid. When omitted, prefix rows
    attend to all valid prefix tokens.
    """
    if action_block_size <= 0:
        raise ValueError("action_block_size must be positive")
    if prefix_len < 0 or action_len < 0:
        raise ValueError("prefix_len and action_len must be non-negative")
    total_len = prefix_len + action_len
    dev = device
    if dev is None:
        for t in (prefix_valid, prefix_att, action_valid):
            if t is not None:
                dev = t.device
                break
    if dev is None:
        dev = torch.device("cpu")

    p_valid = _as_valid(prefix_valid, batch, prefix_len, dev)
    a_valid = _as_valid(action_valid, batch, action_len, dev)

    if prefix_att is None:
        prefix_rows = p_valid[:, :, None] & p_valid[:, None, :]
    else:
        if prefix_att.shape != (batch, prefix_len):
            raise ValueError(
                f"prefix_att must be shaped {(batch, prefix_len)}, got {tuple(prefix_att.shape)}"
            )
        cum = torch.cumsum(prefix_att.to(device=dev, dtype=torch.long), dim=1)
        prefix_rows = (cum[:, None, :] <= cum[:, :, None]) & p_valid[:, :, None] & p_valid[:, None, :]

    prefix_pad = torch.zeros((batch, prefix_len, action_len), dtype=torch.bool, device=dev)
    prefix_rows = torch.cat([prefix_rows, prefix_pad], dim=2)

    nf = prefix_len if non_fast_prefix_len is None else int(non_fast_prefix_len)
    nf = max(0, min(nf, prefix_len))
    action_to_prefix = torch.zeros((batch, action_len, prefix_len), dtype=torch.bool, device=dev)
    if nf > 0:
        action_to_prefix[:, :, :nf] = p_valid[:, None, :nf]
    action_to_prefix &= a_valid[:, :, None]

    q_block = torch.arange(action_len, device=dev) // int(action_block_size)
    kv_block = q_block
    action_block = q_block[:, None] == kv_block[None, :]
    action_block = action_block[None, :, :].expand(batch, -1, -1)
    action_block = action_block & a_valid[:, :, None] & a_valid[:, None, :]
    action_rows = torch.cat([action_to_prefix, action_block], dim=2)

    if prefix_rows.shape != (batch, prefix_len, total_len):
        raise AssertionError("internal prefix mask shape error")
    if action_rows.shape != (batch, action_len, total_len):
        raise AssertionError("internal action mask shape error")
    return prefix_rows, action_rows


def _bool_to_sdpa_mask(mask: torch.Tensor, q: torch.Tensor) -> torch.Tensor:
    value = MASK_VALUE_F32
    if q.dtype.is_floating_point:
        finfo = torch.finfo(q.dtype)
        value = max(MASK_VALUE_F32, finfo.min)
    return torch.where(
        mask[:, None, :, :],
        torch.zeros((), dtype=q.dtype, device=q.device),
        torch.full((), value, dtype=q.dtype, device=q.device),
    )


def _slice_attention_mask(
    attention_mask: torch.Tensor,
    start: int,
    end: int,
    q: torch.Tensor,
) -> torch.Tensor:
    if attention_mask.dim() == 3:
        mask = attention_mask[:, start:end, :]
        if mask.dtype == torch.bool:
            return mask[:, None, :, :]
        return mask[:, None, :, :].to(dtype=q.dtype)
    if attention_mask.dim() == 4:
        mask = attention_mask[:, :, start:end, :]
        return mask if mask.dtype == torch.bool else mask.to(dtype=q.dtype)
    raise ValueError("attention_mask must be (B, S, S) or (B, 1|H, S, S)")


def _sdpa(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    mask: Optional[torch.Tensor],
    *,
    scale: Optional[float],
    dropout_p: float,
    enable_gqa: bool,
) -> torch.Tensor:
    kwargs = {"attn_mask": mask, "dropout_p": float(dropout_p), "scale": scale}
    if enable_gqa:
        kwargs["enable_gqa"] = True
    try:
        return F.scaled_dot_product_attention(q, k, v, **kwargs)
    except TypeError:
        kwargs.pop("enable_gqa", None)
        return F.scaled_dot_product_attention(q, k, v, **kwargs)


def reference_flex_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    *,
    prefix_len: int,
    action_block_size: int,
    attention_mask: Optional[torch.Tensor] = None,
    prefix_valid: Optional[torch.Tensor] = None,
    prefix_att: Optional[torch.Tensor] = None,
    non_fast_prefix_len: Optional[int] = None,
    action_valid: Optional[torch.Tensor] = None,
    detach_prefix_kv_for_action: bool = True,
    scale: Optional[float] = None,
    dropout_p: float = 0.0,
    enable_gqa: Optional[bool] = None,
) -> torch.Tensor:
    """SDPA reference for the PI052 FlexAttention replacement shape.

    Args:
        q, k, v: ``(B, Hq/Hkv, S, D)`` tensors.
        prefix_len: number of prefix rows/columns at the start of sequence.
        action_block_size: size of each block-diagonal action segment.
        attention_mask: optional prebuilt additive or boolean full mask.
        prefix_valid: optional ``(B, P)`` valid prefix positions.
        prefix_att: optional ``(B, P)`` cumulative-block markers.
        non_fast_prefix_len: prefix columns visible to action rows.
        action_valid: optional ``(B, A)`` valid action positions.
        detach_prefix_kv_for_action: detach prefix K/V on the action-row path.
        scale: SDPA scale. Defaults to ``D ** -0.5``.
        dropout_p: SDPA dropout probability.
        enable_gqa: pass SDPA GQA mode when q heads and kv heads differ.
    """
    _check_qkv(q, k, v)
    batch, _, total_len, head_dim = q.shape
    if not (0 <= int(prefix_len) <= total_len):
        raise ValueError("prefix_len must be in [0, S]")
    prefix_len = int(prefix_len)
    action_len = total_len - prefix_len
    if scale is None:
        scale = head_dim**-0.5
    if enable_gqa is None:
        enable_gqa = q.shape[1] != k.shape[1]

    q_prefix = q[:, :, :prefix_len, :]
    q_action = q[:, :, prefix_len:, :]
    k_prefix = k[:, :, :prefix_len, :]
    k_action = k[:, :, prefix_len:, :]
    v_prefix = v[:, :, :prefix_len, :]
    v_action = v[:, :, prefix_len:, :]

    if attention_mask is None:
        prefix_bool, action_bool = build_block_sparse_bool_masks(
            prefix_valid,
            prefix_att,
            batch=batch,
            prefix_len=prefix_len,
            action_len=action_len,
            action_block_size=action_block_size,
            non_fast_prefix_len=non_fast_prefix_len,
            action_valid=action_valid,
            device=q.device,
        )
        prefix_mask = _bool_to_sdpa_mask(prefix_bool, q)
        action_mask = _bool_to_sdpa_mask(action_bool, q)
    else:
        prefix_mask = _slice_attention_mask(attention_mask, 0, prefix_len, q)
        action_mask = _slice_attention_mask(attention_mask, prefix_len, total_len, q)

    out_parts = []
    if prefix_len:
        out_parts.append(
            _sdpa(
                q_prefix,
                k,
                v,
                prefix_mask,
                scale=scale,
                dropout_p=dropout_p,
                enable_gqa=bool(enable_gqa),
            )
        )
    if action_len:
        prefix_k = k_prefix.detach() if detach_prefix_kv_for_action else k_prefix
        prefix_v = v_prefix.detach() if detach_prefix_kv_for_action else v_prefix
        k_for_action = torch.cat([prefix_k, k_action], dim=2)
        v_for_action = torch.cat([prefix_v, v_action], dim=2)
        out_parts.append(
            _sdpa(
                q_action,
                k_for_action,
                v_for_action,
                action_mask,
                scale=scale,
                dropout_p=dropout_p,
                enable_gqa=bool(enable_gqa),
            )
        )
    if not out_parts:
        return q.new_empty(q.shape)
    return torch.cat(out_parts, dim=2) if len(out_parts) == 2 else out_parts[0]


def _manual_attention_part(qs, ks, vs, mask, scale):
    """Materialized-logits attention part: cuBLAS GEMMs + fused masked softmax.

    Same math as SDPA with an additive mask (fp32 softmax; logits stored in
    the io dtype between the GEMM and the softmax). Grouped queries run as a
    strided batched GEMM over the KV heads, so a 1-head K/V is never
    repeated. At PI052 training shapes (GQA 8:1, D=256, bf16) this beats
    both SDPA-with-dense-mask (2.3-3.1x) and the best FlexAttention
    configuration (1.4-2.9x) on fwd+bwd — see benchmarks/RESULTS.md.
    """
    B, H, Sq, D = qs.shape
    Hk = ks.shape[1]
    if Hk != H:
        g = H // Hk
        q2 = qs.reshape(B, Hk, g * Sq, D)
        logits = (q2 @ ks.transpose(-1, -2)).reshape(B, H, Sq, -1)
    else:
        logits = qs @ ks.transpose(-1, -2)
    logits = logits * scale
    if mask is not None:
        logits = logits + mask
    p = logits.float().softmax(dim=-1).to(qs.dtype)
    if Hk != H:
        out = (p.reshape(B, Hk, g * Sq, -1) @ vs).reshape(B, H, Sq, D)
    else:
        out = p @ vs
    return out


# Public alias: integrations (e.g. the LeRobot pi052 flag) consume the raw
# per-part op and assemble masks/splits themselves.
manual_attention_part = _manual_attention_part


def _manual_attention_part_hp(qs, ks, vs, m, scale):
    """High-precision variant: fp32 logits end to end.

    Under torch.compile the ``.float()`` upcasts make the QK product an
    exact fp32 GEMM, removing the bf16 rounding of the logits that
    dominates the default variant's error (softmax-output max-abs error
    drops ~16x, 9.8e-4 -> 6.1e-5 at PI052 shapes). Costs roughly 3x on
    the QK+softmax stage (~5-6 ms per training step at B=2) because the
    fp32 GEMM does not use the bf16 tensor-core path — use where parity
    matters more than the last few percent of speed.
    """
    B, H, Sq, D = qs.shape
    Hk = ks.shape[1]
    if Hk != H:
        g = H // Hk
        q2 = qs.reshape(B, Hk, g * Sq, D)
        logits = (q2.float() @ ks.transpose(-1, -2).float()).reshape(B, H, Sq, -1)
    else:
        logits = qs.float() @ ks.transpose(-1, -2).float()
    logits = logits * scale
    if m is not None:
        logits = logits + m.float()
    p = logits.softmax(dim=-1).to(qs.dtype)
    if Hk != H:
        out = (p.reshape(B, Hk, g * Sq, -1) @ vs).reshape(B, H, Sq, D)
    else:
        out = p @ vs
    return out


manual_attention_part_hp = _manual_attention_part_hp


def _softmax_bwd_chain(p, dp, scale):
    p32 = p.float()
    dp32 = dp.float()
    return (p32 * (dp32 - (dp32 * p32).sum(dim=-1, keepdim=True)) * scale).to(p.dtype)


_softmax_bwd_compiled = None


def _get_softmax_bwd():
    global _softmax_bwd_compiled
    if _softmax_bwd_compiled is None:
        _softmax_bwd_compiled = torch.compile(_softmax_bwd_chain, dynamic=False)
    return _softmax_bwd_compiled


class _ManualAttentionPartFn(torch.autograd.Function):
    """Manual attention part with bf16-saved probabilities.

    Same math as :func:`_manual_attention_part`; the backward is written
    out so only the io-dtype probability tensor is saved (autograd on the
    composed version keeps the fp32 softmax output alive — 3x the bytes).
    The softmax gradient itself is still computed in fp32.
    """

    @staticmethod
    def forward(ctx, q, k, v, mask, scale):
        B, H, Sq, D = q.shape
        Hk = k.shape[1]
        if Hk != H:
            g = H // Hk
            q2 = q.reshape(B, Hk, g * Sq, D)
            logits = (q2 @ k.transpose(-1, -2)).reshape(B, H, Sq, -1)
        else:
            logits = q @ k.transpose(-1, -2)
        logits = logits * scale
        if mask is not None:
            logits = logits + mask
        p = logits.float().softmax(dim=-1).to(q.dtype)
        if Hk != H:
            out = (p.reshape(B, Hk, g * Sq, -1) @ v).reshape(B, H, Sq, D)
        else:
            out = p @ v
        ctx.save_for_backward(q, k, v, p)
        ctx.scale = scale
        return out

    @staticmethod
    def backward(ctx, dout):
        q, k, v, p = ctx.saved_tensors
        scale = ctx.scale
        B, H, Sq, D = q.shape
        Hk = k.shape[1]
        dout = dout.contiguous()
        if Hk != H:
            g = H // Hk
            dout2 = dout.reshape(B, Hk, g * Sq, D)
            p2 = p.reshape(B, Hk, g * Sq, -1)
            dp = (dout2 @ v.transpose(-1, -2)).reshape(B, H, Sq, -1)
            dv = p2.transpose(-1, -2) @ dout2
        else:
            dp = dout @ v.transpose(-1, -2)
            dv = p.transpose(-1, -2) @ dout
        ds = _get_softmax_bwd()(p, dp, scale)
        if Hk != H:
            ds2 = ds.reshape(B, Hk, g * Sq, -1)
            dq = (ds2 @ k).reshape(B, H, Sq, D)
            dk = ds2.transpose(-1, -2) @ q.reshape(B, Hk, g * Sq, D)
        else:
            dq = ds @ k
            dk = ds.transpose(-1, -2) @ q
        return dq, dk, dv, None, None


def manual_attention_part_v2(q, k, v, mask, scale):
    """bf16-saved-p variant of :func:`manual_attention_part` (fwd+bwd)."""
    return _ManualAttentionPartFn.apply(q, k, v, mask, scale)

_manual_part_compiled = None


def _get_manual_part():
    global _manual_part_compiled
    if _manual_part_compiled is None:
        _manual_part_compiled = torch.compile(_manual_attention_part, dynamic=False)
    return _manual_part_compiled


def manual_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    *,
    prefix_len: int,
    action_block_size: int,
    attention_mask: Optional[torch.Tensor] = None,
    prefix_valid: Optional[torch.Tensor] = None,
    prefix_att: Optional[torch.Tensor] = None,
    non_fast_prefix_len: Optional[int] = None,
    action_valid: Optional[torch.Tensor] = None,
    detach_prefix_kv_for_action: bool = True,
    scale: Optional[float] = None,
    dropout_p: float = 0.0,
    compile_part: bool = True,
) -> torch.Tensor:
    """Materialized-logits implementation of :func:`reference_flex_attention`.

    Same mask semantics and prefix/action split; each part runs through
    :func:`_manual_attention_part` instead of SDPA. ``dropout_p`` must be 0
    (training attention dropout is unused in PI052); other values raise so
    callers fall back explicitly.
    """
    if dropout_p:
        raise ValueError("manual_attention does not support dropout; use the reference path")
    _check_qkv(q, k, v)
    batch, _, total_len, head_dim = q.shape
    if not (0 <= int(prefix_len) <= total_len):
        raise ValueError("prefix_len must be in [0, S]")
    prefix_len = int(prefix_len)
    action_len = total_len - prefix_len
    if scale is None:
        scale = head_dim**-0.5

    if attention_mask is None:
        prefix_bool, action_bool = build_block_sparse_bool_masks(
            prefix_valid,
            prefix_att,
            batch=batch,
            prefix_len=prefix_len,
            action_len=action_len,
            action_block_size=action_block_size,
            non_fast_prefix_len=non_fast_prefix_len,
            action_valid=action_valid,
            device=q.device,
        )
        prefix_mask = _bool_to_sdpa_mask(prefix_bool, q)
        action_mask = _bool_to_sdpa_mask(action_bool, q)
    else:
        prefix_mask = _slice_attention_mask(attention_mask, 0, prefix_len, q)
        action_mask = _slice_attention_mask(attention_mask, prefix_len, total_len, q)
        if prefix_mask.dtype == torch.bool:
            prefix_mask = _bool_to_sdpa_mask(prefix_mask[:, 0], q)
        if action_mask.dtype == torch.bool:
            action_mask = _bool_to_sdpa_mask(action_mask[:, 0], q)

    part = _get_manual_part() if compile_part else _manual_attention_part
    out_parts = []
    if prefix_len:
        out_parts.append(part(q[:, :, :prefix_len, :], k, v, prefix_mask, scale))
    if action_len:
        k_prefix = k[:, :, :prefix_len, :]
        v_prefix = v[:, :, :prefix_len, :]
        if detach_prefix_kv_for_action:
            k_prefix = k_prefix.detach()
            v_prefix = v_prefix.detach()
        k_for_action = torch.cat([k_prefix, k[:, :, prefix_len:, :]], dim=2)
        v_for_action = torch.cat([v_prefix, v[:, :, prefix_len:, :]], dim=2)
        out_parts.append(part(q[:, :, prefix_len:, :], k_for_action, v_for_action, action_mask, scale))
    if not out_parts:
        return q.new_empty(q.shape)
    return torch.cat(out_parts, dim=2) if len(out_parts) == 2 else out_parts[0]


def flex_attention(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    *,
    prefix_len: int,
    action_block_size: int,
    attention_mask: Optional[torch.Tensor] = None,
    prefix_valid: Optional[torch.Tensor] = None,
    prefix_att: Optional[torch.Tensor] = None,
    non_fast_prefix_len: Optional[int] = None,
    action_valid: Optional[torch.Tensor] = None,
    detach_prefix_kv_for_action: bool = True,
    scale: Optional[float] = None,
    dropout_p: float = 0.0,
    enable_gqa: Optional[bool] = None,
    force_fallback: bool = False,
    impl: str = "sdpa",
) -> torch.Tensor:
    """Flex-style block-sparse attention.

    ``impl="sdpa"`` (default) keeps the SDPA reference path;
    ``impl="manual"`` routes through the materialized-logits
    implementation; ``impl="auto"`` picks manual only where it has been
    measured to win end-to-end — consumer Blackwell (sm120-class) with
    no dropout. On A100 (sm80) the manual math wins microbenches but
    loses training-step integration, and on H100/H200 (sm90) the fused
    FMHA kernels win outright, so auto keeps SDPA there.
    """
    _ = force_fallback
    if impl == "auto":
        sm120 = q.is_cuda and torch.cuda.get_device_capability(q.device)[0] == 12
        impl = "manual" if (sm120 and not dropout_p) else "sdpa"
    if impl == "manual":
        return manual_attention(
            q,
            k,
            v,
            prefix_len=prefix_len,
            action_block_size=action_block_size,
            attention_mask=attention_mask,
            prefix_valid=prefix_valid,
            prefix_att=prefix_att,
            non_fast_prefix_len=non_fast_prefix_len,
            action_valid=action_valid,
            detach_prefix_kv_for_action=detach_prefix_kv_for_action,
            scale=scale,
            dropout_p=dropout_p,
        )
    return reference_flex_attention(
        q,
        k,
        v,
        prefix_len=prefix_len,
        action_block_size=action_block_size,
        attention_mask=attention_mask,
        prefix_valid=prefix_valid,
        prefix_att=prefix_att,
        non_fast_prefix_len=non_fast_prefix_len,
        action_valid=action_valid,
        detach_prefix_kv_for_action=detach_prefix_kv_for_action,
        scale=scale,
        dropout_p=dropout_p,
        enable_gqa=enable_gqa,
    )


def flex_attention_forward(*args, **kwargs) -> torch.Tensor:
    """Forward-only compatibility wrapper."""
    return flex_attention(*args, **kwargs)


def backend_marker(x: torch.Tensor) -> torch.Tensor:
    if ops is None:
        return x
    return ops._flashrt_training_package_marker(x)


__all__ = [
    "MASK_VALUE_F32",
    "backend_marker",
    "build_block_sparse_bool_masks",
    "flex_attention",
    "flex_attention_forward",
    "manual_attention",
    "manual_attention_part",
    "manual_attention_part_hp",
    "manual_attention_part_v2",
    "reference_flex_attention",
]