Skip to main content

DisentangledFlash

PyPI Version CI License

Fast exact DeBERTa-style disentangled attention in Triton.

DisentangledFlash provides fused Triton and optimized PyTorch inference backends for bidirectional DeBERTa-v2/v3 attention without materializing the [B, H, L, L] attention matrix.

Status

  • NVIDIA CUDA + Triton
  • inference only
  • DeBERTa-v2/v3-style C2P/P2C disentangled attention
  • FP16, BF16, strict FP32, and optional fast FP32/TF32 mode
  • head dimensions 32, 64, and 128
  • factorized 2-D padding masks
  • runtime sequence lengths through nine bounded kernel families up to 8192
  • FlashAttention-style packed, unpadded inference with cu_seqlens
  • fused QKV projection is always enabled for the PyTorch/Triton backends

Training/backward is not implemented. CPU/MPS use the PyTorch backend for development/benchmarking, not the Triton kernel.

Install

pip install -e ".[dev]"

For the pretrained GLUE/MNLI evaluation:

pip install -e ".[dev,benchmark,hf]"

Hugging Face usage

from transformers import AutoModelForSequenceClassification
from disentangled_flash import optimize_deberta

model = (
    AutoModelForSequenceClassification.from_pretrained("microsoft/deberta-v2-xlarge-mnli")
    .cuda()
    .eval()
)

optimize_deberta(
    model.deberta,
    sequence_lengths=[64, 128, 384, 512, 768, 1024, 2048, 4096, 8192],
)

Only the encoder attention path is replaced; checkpoint names and the remaining model layers are preserved. Use enable_deberta_inference(..., backend="torch") to select the PyTorch backend.

CUDA benchmark

pip install -e ".[benchmark]"
python -m benchmarks.benchmark_cuda \
  --output deberta_v3_base_encoder_cuda_results.json

The default matrix compares the full DeBERTa-v3-base encoder across Hugging Face, DF PyTorch, DF Triton, and FlashDeBERTa at lengths 64–8192 in padded and packed modes. It uses fresh deterministic inputs, records OOMs without stopping, and saves hardware, software, Slurm, git, command, and environment metadata.

Generate latency and memory plots with:

python -m benchmarks.plot_cuda_results \
  deberta_v3_base_encoder_cuda_results.json \
  --output-dir benchmarks/results/deberta_v3_base_encoder

Kernel tuning profiles

Matching saved profiles are used automatically; otherwise Triton runs bounded autotuning on first use. Override this with:

from disentangled_flash import KernelConfig, KernelTuningOptions, optimize_deberta

optimize_deberta(model, tuning=KernelTuningOptions(mode="autotune"))
optimize_deberta(
    model,
    tuning=KernelTuningOptions(profile_paths=("my-gpu-profile.json",)),
)
optimize_deberta(
    model,
    tuning=KernelTuningOptions(
        mode="fixed",
        fixed_config=KernelConfig(64, 64, 4),
    ),
)

Profiles are discovered in the user cache, DISENTANGLED_FLASH_PROFILE_DIR, and the package. Compatibility is keyed by GPU/compiler stack and workload; the driver is diagnostic only. Exact length, batch size, head count, and active-slot count are runtime values rather than autotune keys.

Generate a resumable profile on a CUDA machine with:

python -m disentangled_flash.tune \
  --preset standard \
  --output rtx-6000-ada.json
python -m disentangled_flash.tune inspect rtx-6000-ada.json

The standard preset covers length families 64, 128, 384, 512, 768, 1024, 2048, 4096, and 8192, supported dtypes, attention modes, occupancy regimes, and padded/packed layouts. Results are parity-checked and saved after each workload so tuning can resume.

Packed unpadded inference

forward_packed(hidden_states, cu_seqlens, max_seqlen) accepts FlashAttention-style [total_tokens, hidden_size] input. Convert right-padded input with pack_padded_with_info and reuse its metadata across encoder layers:

from disentangled_flash import pack_padded_with_info, unpack_packed

tokens, cu_seqlens, info = pack_padded_with_info(hidden_states, attention_mask)
packed_output = encoder.forward_packed(
    tokens,
    cu_seqlens,
    info.max_seqlen,
    packed_info=info,
).last_hidden_state
output, _ = unpack_packed(
    packed_output,
    cu_seqlens,
    hidden_states.size(1),
    packed_info=info,
)

Pretrained GLUE/MNLI evaluation

python -m benchmarks.evaluate_mnli

This runs one cold, full GLUE/MNLI matched-validation pass at batch size 8 across Hugging Face, DF PyTorch, DF Triton, and FlashDeBERTa. It reports performance, accuracy, numerical error, and classification-decision parity.

To require the bundled H200 configuration and prohibit Triton autotuning:

python -m benchmarks.evaluate_mnli \
  --tuning-mode profile_only \
  --profile src/disentangled_flash/profiles/h200-sm90-deberta-v3-base-torch-2.14-cu130-triton-3.8.json

CUDA validation

python -m validation.validate_cuda

Run this before publishing a profile for a new GPU or compiler stack.

Results: H200 DeBERTa-v3-base encoder

H200 DeBERTa-v3-base latency at batch 1

H200 DeBERTa-v3-base latency at batch 16

H200 DeBERTa-v3-base peak memory at batch 1

H200 DeBERTa-v3-base peak memory at batch 16

These H200 results cover the complete 12-layer DeBERTa-v3-base encoder and report batch sizes 1 and 16.

Benchmark configuration

Parameter Value
Model microsoft/deberta-v3-base architecture
Reported batch sizes 1, 16
Sequence lengths 64, 128, 256, 512, 1024, 2048, 4096, 8192
Precisions FP16, BF16, strict FP32
Execution / measurements Eager; 3 warmups and 10 fresh inputs per point
Packed-length distribution Uniform from 60% through 100% of the padded length
GPU NVIDIA H200, SM 9.0, 143771 MiB VRAM
Driver / power limit 595.91.07 / 700 W
Software Python 3.12.14, PyTorch 2.14.0+cu130, Triton 3.8.0, cuDNN 9.2.4
Comparisons Transformers 5.17.0, FlashDeBERTa 0.0.7
Host allocation 16 CPU threads and 128 GiB RAM under Slurm
Host / OS Xeon Platinum 8568Y+; Linux 6.18.51-1-insait, x86-64, glibc 2.41

Latency is the sample mean and excludes preparation, compilation, and offline tuning. Peak memory is total CUDA allocation. Corresponding implementations use the same deterministic samples; OOM points are capacity results.

Latency

Geometric-mean packed-Triton speedups across the eight sequence lengths:

Batch Precision vs. Hugging Face padded vs. DF PyTorch packed vs. FlashDeBERTa packed
1 FP16 1.66× 2.07× 1.22×
1 BF16 1.71× 2.12× 1.23×
1 FP32 1.52× 1.43× 1.24×
16 FP16 2.33× 5.75× 1.45×
16 BF16 2.29× 5.62× 1.38×
16 FP32 1.51× 2.33× 1.18×

Batch-16 Hugging Face and DF PyTorch comparisons stop at 4096 because both OOM at 8192. The following table compares successful length-8192 points:

Batch Precision DF Triton packed FlashDeBERTa packed Speedup DF Triton peak FlashDeBERTa peak Memory reduction
1 FP16 40.13 ms 72.13 ms 1.80× 1.01 GiB 3.07 GiB 67.1%
1 BF16 39.35 ms 68.67 ms 1.75× 1.01 GiB 3.07 GiB 67.1%
1 FP32 155.66 ms 191.80 ms 1.23× 1.99 GiB 3.57 GiB 44.4%
16 FP16 726.94 ms 1093.03 ms 1.50× 5.13 GiB 17.39 GiB 70.5%
16 BF16 708.29 ms 1034.98 ms 1.46× 5.13 GiB 17.39 GiB 70.5%
16 FP32 2921.89 ms 3116.86 ms 1.07× 10.24 GiB 26.22 GiB 60.9%

Packing conversion has a fixed cost; at batch 16 FP16/BF16, packed Triton overtakes padded Triton at length 512.

Peak allocated GPU memory

At batch 16 and length 8192, Hugging Face and packed DF PyTorch OOM in every precision; padded DF PyTorch OOMs in FP32. Both Triton layouts and FlashDeBERTa complete all three precisions.

Bundled H200 tuning profile

The package automatically discovers the reviewed h200-sm90-deberta-v3-base-torch-2.14-cu130-triton-3.8.json profile. Its 108 winners cover the DeBERTa-v3-base layouts, precisions, and bounded lengths through 8192. They require H200 SM 9.0, PyTorch 2.14.0+cu130, CUDA 13.0, and the recorded Triton 3.8.0 compiler fingerprint; otherwise auto mode uses bounded autotuning.

Pretrained-model parity

The H200 task-level test uses microsoft/deberta-v2-xlarge-mnli, FP16, batch 16, length 512, and one cold pass over all 9,815 matched-validation examples. The dense input is 92.62% padding. Every implementation reaches 91.7371% accuracy with full decision parity (0/9,815 mismatches).

Implementation Time Throughput Speedup vs. HF Decision mismatches Logit abs. error max / mean Probability abs. error max / mean
Hugging Face padded 55,963.399 ms 175.38 examples/s 1.00× 0 / 9,815 0 / 0 0 / 0
DF PyTorch padded 53,484.569 ms 183.51 examples/s 1.05× 0 / 9,815 0.109375 / 0.00126508 0.0100614 / 0.00008926
DF Triton padded 31,329.634 ms 313.28 examples/s 1.79× 0 / 9,815 0.0742188 / 0.00114972 0.00796831 / 0.00008365
DF Triton packed 5,419.671 ms 1,811.00 examples/s 10.33× 0 / 9,815 0.0507812 / 0.00113259 0.00853068 / 0.00008150
FlashDeBERTa packed 22,058.000 ms 444.96 examples/s 2.54× 0 / 9,815 0.0800781 / 0.00115263 0.00675502 / 0.00008179

Packed Triton is 4.07× faster than FlashDeBERTa here. Error columns report full-dataset maximum and mean absolute error; one pass provides no run-to-run standard deviation.

Attribution

The auditable reference implementation is derived from Hugging Face Transformers 4.57.6 DeBERTa-v2/v3 modeling code and retains its original Apache-2.0 header. See THIRD_PARTY_NOTICES.md.

Citation

If you use DisentangledFlash in your research or project, please cite it as follows:

@software{boychev2026disentangledflash,
  author = {Boychev, Delyan},
  title = {DisentangledFlash: Fast exact DeBERTa-style disentangled attention in Triton},
  url = {https://github.com/delyan-boychev/disentangled-flash},
  version = {0.2.0},
  year = {2026}
}

Release files for disentangled-flash 0.2.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for disentangled-flash 0.2.0
File Size Uploaded
disentangled_flash-0.2.0.tar.gz 68.8 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for disentangled-flash 0.2.0
File Interpreter ABI Platform
disentangled_flash-0.2.0-py3-none-any.whl Python 3 none any Details

Total release size: 129.2 kB

Release files / disentangled_flash-0.2.0.tar.gz

Download URL disentangled_flash-0.2.0.tar.gz
Size 68.8 kB
Tags Source
SHA-256 checksum
How to use checksums
8301f256b646eb115e35620471bde8aaf1085f87a01e140fb28d8cfb27d46706
BLAKE2b-256 checksum
How to use checksums
6439da464928559361480ebce334fd4f4debfd973a6a2138697bd813d12eaa1a
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 17, 2026.

Transparency log

Release files / disentangled_flash-0.2.0-py3-none-any.whl

Download URL disentangled_flash-0.2.0-py3-none-any.whl
Size 60.4 kB
Tags Python 3
SHA-256 checksum
How to use checksums
d84749260ddb3ddd7ff9019d52d211069546403f932f910107c0e29390a4a0a7
BLAKE2b-256 checksum
How to use checksums
aadefe7247c0f3777ac56595b5a433515d52a4b7bf34a8c4c9af87cbfc1d092f
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 17, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

0.2.0 This release

2 release files

0.1.4

2 release files

0.1.3

2 release files

0.1.2

2 release files

0.1.1

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page