Metadata-Version: 2.4
Name: fast-attnres
Version: 2.0.1
Summary: CUDA BF16 PyTorch operators for standard and sliced low-rank Attention Residuals
Author: Jonathan Su
License: MIT
Project-URL: Homepage, https://github.com/jon123boss/fast-attnres
Project-URL: Repository, https://github.com/jon123boss/fast-attnres
Project-URL: Issues, https://github.com/jon123boss/fast-attnres/issues
Project-URL: Documentation, https://github.com/jon123boss/fast-attnres#readme
Project-URL: Citation, https://github.com/jon123boss/fast-attnres/blob/v2.0.1/CITATION.cff
Project-URL: Paper, https://arxiv.org/abs/2603.15031
Project-URL: LR-AttnRes Paper, https://arxiv.org/abs/2607.09694
Keywords: pytorch,attention,residuals,triton,cuda,bfloat16
Classifier: Development Status :: 4 - Beta
Classifier: Intended Audience :: Developers
Classifier: Programming Language :: Python :: 3
Classifier: Programming Language :: Python :: 3 :: Only
Classifier: Programming Language :: Python :: 3.10
Classifier: Programming Language :: Python :: 3.11
Classifier: Typing :: Typed
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Requires-Python: >=3.10
Description-Content-Type: text/markdown
License-File: LICENSE
License-File: NOTICE
Requires-Dist: torch>=2.9
Provides-Extra: cuda
Requires-Dist: torch==2.13.0; (platform_system == "Linux" and platform_machine == "x86_64") and extra == "cuda"
Requires-Dist: triton==3.7.1; (platform_system == "Linux" and platform_machine == "x86_64") and extra == "cuda"
Provides-Extra: test
Requires-Dist: pytest>=7; extra == "test"
Requires-Dist: numpy<3,>=1.24; extra == "test"
Requires-Dist: tomli>=2; python_version < "3.11" and extra == "test"
Provides-Extra: dev
Requires-Dist: build>=1; extra == "dev"
Requires-Dist: ruff>=0.8; extra == "dev"
Provides-Extra: benchmark
Requires-Dist: numpy<3,>=1.24; extra == "benchmark"
Requires-Dist: einops>=0.8; extra == "benchmark"
Requires-Dist: modal==1.5.4; extra == "benchmark"
Provides-Extra: plot
Requires-Dist: matplotlib<4,>=3.8; extra == "plot"
Requires-Dist: numpy<3,>=1.24; extra == "plot"
Dynamic: license-file

# Fast Attention Residuals

[![CI](https://github.com/jon123boss/fast-attnres/actions/workflows/ci.yml/badge.svg)](https://github.com/jon123boss/fast-attnres/actions/workflows/ci.yml)
[![Python 3.10+](https://img.shields.io/badge/python-3.10%2B-3776AB.svg)](https://www.python.org/)
[![PyTorch 2.13](https://img.shields.io/badge/tested-PyTorch_2.13-EE4C2C.svg)](https://pytorch.org/)
[![Triton 3.7.1](https://img.shields.io/badge/tested-Triton_3.7.1-654FF0.svg)](https://github.com/triton-lang/triton/releases/tag/v3.7.1)
[![License: MIT](https://img.shields.io/badge/license-MIT-2E7D32.svg)](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/LICENSE)

![Full AttnRes training on H100 SXM and B200](https://raw.githubusercontent.com/jon123boss/fast-attnres/v2.0.1/results/final_sweep/compiled_step_hero.png)

**Fast Attention Residuals** (`Fast-AttnRes`) makes
[Attention Residuals](https://arxiv.org/abs/2603.15031) a single PyTorch
operation: pass ordered full-width residual sources and one learned query, get
one full-width residual back. The same `attnres(values, query)` call handles
standard and sliced low-rank AttnRes in Full and Block schedules, with packed
tensors or ordered source lists.

Start with the [standard quickstart](#quickstart-standard-attnres), choose a
[Full or Block schedule](#full-and-block-schedules), and try
[sliced LR-AttnRes](#sliced-lr-attnres) for a smaller routing query.

## Why use Fast-AttnRes

- **One PyTorch call:** full-width output and ordinary first-order autograd,
  including gradients through routing keys and shared residual sources.
- **One schedule primitive:** Full and Block use the same operator; the caller
  controls source order, block sums, and learned queries.
- **BF16 training:** CUDA BF16 values and queries, FP32 internal reductions,
  `torch.compile`, and CUDA Graph replay on H100 and B200.

## Training performance

On the 24-layer Full workload, Fast-AttnRes reduces step latency by **5.95% on
H100 SXM** and **22.74% on B200** against native FLA Triton checkpoint 1
(median of three paired seed estimates).

The figures measure complete BF16 CUDA Graph training steps on H100 SXM and
B200, including forward, backward, and the optimizer update. Full and Block
use the same per-read Fast-AttnRes backend.

The 24-layer headline uses three seeds with 120 paired rounds each. The
8-layer comparisons use one seed with 40 paired rounds per configuration.
See the [benchmark protocol](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/docs/benchmark_results.md) and
[reproducible results](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/results/final_sweep/README.md) for workloads, source
versions, confidence intervals, and reproduction commands.

### Compiled BF16 training steps

Comparators are native FLA Triton checkpoint 1, Liger 0.8.2, and Catswe phase 1.
Unsupported and failed arms remain labelled. Quarter-rank comparisons against
standard FLA compare different routing equations.
The D2048 results retain [disclosed normalization-rounding deviations](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/docs/benchmark_results.md).
On H100, standard D2048 is 0.28% slower than FLA, within the declared 1% parity band.

![H100 compiled BF16 training steps](https://raw.githubusercontent.com/jon123boss/fast-attnres/v2.0.1/results/final_sweep/compiled_step_sweep_h100.png)

![B200 compiled BF16 training steps](https://raw.githubusercontent.com/jon123boss/fast-attnres/v2.0.1/results/final_sweep/compiled_step_sweep_b200.png)

### Quarter-rank routing

These figures compare our `R=D/4` kernel with our `R=D` kernel on each workload.
Values and outputs retain width `D`; only the routing rank changes.
Quarter rank reduces step latency in all ten measured device/workload pairs.

![H100 quarter-rank versus full-rank routing](https://raw.githubusercontent.com/jon123boss/fast-attnres/v2.0.1/results/final_sweep/rank_comparison_h100.png)

![B200 quarter-rank versus full-rank routing](https://raw.githubusercontent.com/jon123boss/fast-attnres/v2.0.1/results/final_sweep/rank_comparison_b200.png)

## Install

Install the release with the pinned CUDA runtime:

```bash
python -m pip install --index-url https://download.pytorch.org/whl/cu130 torch==2.13.0
python -m pip install "fast-attnres[cuda]==2.0.1"
```

For development, clone the repository and run `python -m pip install -e ".[cuda]"`.

Version 2 requires CUDA BF16 tensors; the CPU/FP32 execution and exported
reference from version 1 are no longer part of the public API.

The pinned runtime is Python 3.11, PyTorch 2.13.0 with CUDA 13.0, and Triton 3.7.1.

## Quickstart: standard AttnRes

```python
import torch
from attnres import attnres

values = torch.randn(8, 2, 1024, device="cuda", dtype=torch.bfloat16, requires_grad=True)
source_list = tuple(values.unbind(0))
query = torch.randn(1024, device="cuda", dtype=torch.bfloat16, requires_grad=True)
read = torch.compile(attnres, fullgraph=True, dynamic=False)
output = read(source_list, query)
output.square().mean().backward()
```

The public signature is `attnres(values, query, *, eps=2**-23, scale=1.0)`.
`values` is either packed `[S, ..., D]` or an ordered list/tuple of `[ ..., D ]`
tensors. The query is `[R]`, with `1 <= R <= D`; the output retains width `D`.
Values, query, output, and first-order operator gradients are CUDA BF16. Internal
FP32 accumulators may be used for normalization, logits, softmax, and reductions.

## Equation

For source value `v_s`, take its final `R` coordinates as the implicit key:

```text
t_s       = v_s[..., D-R:D]
r_s       = sqrt(mean(t_s ** 2) + eps)
k_s       = t_s / r_s
score_s   = scale * dot(k_s, query)
p_s       = softmax(score, axis=source)_s
output    = sum_s p_s * v_s
```

Normalization and softmax run independently at each carried batch or token
position. The output is not normalized, source-count weighted, or source averaged.
See [`docs/equation.md`](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/docs/equation.md) for the complete contract.

## Full and Block schedules

Full supplies the embedding and every preceding writer output. Block supplies
the embedding, completed block sums, and an optional current partial sum. The
caller owns block boundaries and sums a partial block before the read; it is
passed as one ordinary source and is not averaged.

```python
full_output = attnres((embedding, *writers), query)

completed = (embedding, first_block, second_block)
block_sources = completed if partial is None else completed + (partial,)
block_output = attnres(block_sources, query)
```

Every read evaluates its routing weights and mixture from the supplied sources.
Both schedules use the same query, `eps`, `scale`, and operator contract.

## Sliced LR-AttnRes

Sliced LR-AttnRes keeps full-width values and output while using a shorter
implicit key and query:

```python
rank = 64
query = torch.randn(rank, device="cuda", dtype=torch.bfloat16)
output = attnres(source_list, query)  # [..., D], BF16
```

For a trainable static query, use `LearnedQuery(rank)`:

```python
from attnres import LearnedQuery
learned_query = LearnedQuery(rank).to(device="cuda", dtype=torch.bfloat16)
output = attnres(source_list, learned_query())
```

Standard AttnRes is the `R == D` case; sliced routing uses `R < D` and the final
`R` value coordinates. Projected keys, routing priors, and architectural changes
are outside this package contract.

Optimization targets common ranks: powers of two from 16 upward, plus widths
such as 384 and 640. Other ranks retain the same mathematical support.

## Validation scope

Correctness checks compare BF16 outputs and first-order gradients with an
independent BF16 PyTorch reference with FP32 internal accumulation at
`rtol=0.05` and `atol=0.05`. Coverage
includes packed/list sources, repeated reads, partial Blocks, changed inputs,
non-contiguous layouts, shared sources, and compiled replay. See the
[validation protocol](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/docs/validation.md).

Timing reports identify the exact source, device, runtime, workload, and
measurement boundary. Failed, incomplete, and inconclusive comparisons remain
visible in the results.

## License

Fast-AttnRes is released under the [MIT License](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/LICENSE). FLA-derived source
list attribution remains in [`NOTICE`](https://github.com/jon123boss/fast-attnres/blob/v2.0.1/NOTICE).
