Spaces:
Running
Fused Triton attention for DeBERTa-v2/v3 - Disentangled Flash
Hey kernels-community,
I've been working on a fused Triton implementation of DeBERTa-v2/v3 disentangled attention, and I'd like publishing access for delyanboychev so I can publish it as delyanboychev/disentangled-flash at v0.
Motivation
DeBERTa-v2/v3 cannot use standard SDPA or FlashAttention directly because its disentangled attention adds content-to-position and position-to-content terms. As a result, the reference encoder materializes the full [B, H, L, L] attention matrix.
DisentangledFlash computes the same attention operation using a tiled, streaming Triton kernel without materializing that matrix.
There is no approximation involved: it evaluates the same attention formulation with a different execution strategy.
End-to-end MNLI result
On microsoft/deberta-v2-xlarge-mnli, FP16, batch size 16, sequence length 512, running one cold pass over all 9,815 MNLI matched-validation examples:
The input is 92.6% padding.
| Implementation | Time | Throughput | Speedup | Different predictions |
|---|---|---|---|---|
| Transformers | 55,963 ms | 175.38 ex/s | 1.00× | — |
| This kernel (padded) | 31,330 ms | 313.28 ex/s | 1.79× | 0 / 9,815 |
| FlashDeBERTa (packed) | 22,058 ms | 444.96 ex/s | 2.54× | 0 / 9,815 |
| This kernel (packed) | 5,420 ms | 1,811.00 ex/s | 10.33× | 0 / 9,815 |
All paths reach exactly 91.7371% accuracy, with zero differing predictions.
DeBERTa-v3-base encoder benchmarks
On H200, evaluated across sequence lengths 64–8192, the geometric-mean speedups are:
| Batch | Precision | vs. Transformers | vs. FlashDeBERTa |
|---|---|---|---|
| 1 | FP16 | 1.66× | 1.22× |
| 1 | BF16 | 1.71× | 1.23× |
| 1 | FP32 | 1.52× | 1.24× |
| 16 | FP16 | 2.33× | 1.45× |
| 16 | BF16 | 2.29× | 1.38× |
| 16 | FP32 | 1.51× | 1.18× |
Long-context memory usage
The memory reduction is the part I find most useful.
At sequence length 8192, where the Transformers reference implementation OOMs at batch size 16:
| Batch | Precision | Ours | FlashDeBERTa | Speedup | Our peak | Their peak | Memory saved |
|---|---|---|---|---|---|---|---|
| 1 | FP16 | 40.13 ms | 72.13 ms | 1.80× | 1.01 GiB | 3.07 GiB | 67% |
| 1 | BF16 | 39.35 ms | 68.67 ms | 1.75× | 1.01 GiB | 3.07 GiB | 67% |
| 1 | FP32 | 155.66 ms | 191.80 ms | 1.23× | 1.99 GiB | 3.57 GiB | 44% |
| 16 | FP16 | 726.94 ms | 1093.03 ms | 1.50× | 5.13 GiB | 17.39 GiB | 70% |
| 16 | BF16 | 708.29 ms | 1034.98 ms | 1.46× | 5.13 GiB | 17.39 GiB | 70% |
| 16 | FP32 | 2921.89 ms | 3116.86 ms | 1.07× | 10.24 GiB | 26.22 GiB | 61% |
One important caveat: the kernel runs in two modes, and the encoder geometric means and the length-8192 table above are both packed.
Padded takes the [B, L, H] batch that the current Transformers encoder interface hands you, and that is what the Hub layer exposes today. It corresponds to the 1.79× MNLI result above.
Packed concatenates the batch into a single [total_tokens, H] sequence with cu_seqlens and never computes on padding at all, which is where the larger numbers come from.
Same kernel, same math; the difference is the input layout and the padding overhead.
GPU support
Turing has a substantially different register-pressure and occupancy tradeoff, while the current tile configurations were tuned for Ada and Hopper.
Because of that, I would prefer the kernel to be mapped only to compute capabilities on which it has been validated rather than enabling it for every CUDA GPU by default.
Packaging and validation
The package is:
- pure Triton
- edition 5
[torch-noarch]- CUDA
- exposes
layers.DebertaV2Encoder
check-config and the full build pass in CI.
The get_kernel import check also passes with the Triton autotune decorators enabled, so there is no doGetKernelCheck override.
Validation currently includes:
- 24 CPU tests covering the Hub layer-contract requirements
- 39 tests on an RTX 6000 Ada
- FP16, BF16, and FP32 parity against the reference encoder
- head dimensions 64 and 128
- fallback on unsupported shapes
- heavy-padding cases with no non-finite outputs
- training-mode fallback to the reference forward path (i am working on the backward, soon it will be ready)
The layer is stateless and declares:
has_backward = False
It does not replace encoder submodules, copy weights, or modify checkpoint keys.
Instead, it shares the host encoder's projections by reference and replaces only the attention computation. Unsupported cases transparently fall back to the stock Transformers implementation.
Source / Transformers integration
Sources:
https://github.com/delyan-boychev/disentangled-flash
https://github.com/delyan-boychev/disentangled-flash-kernel
The kernel is intended to back the DebertaV2Encoder kernel hook in:
huggingface/transformers#48176
That integration is ready, but it pins a revision of this kernel, so publishing the Hub package is currently the blocker.
I'm happy to adjust the package layout, naming, supported-device mapping, or anything else needed before publishing.
Mmm