feat(backend): sliding-window PagedAttention #116

Merged
rcsheets merged 1 commit from feat/sliding-window-attention into main 2026-10-04 06:23:27 +00:00
Owner

backend.Ops.PagedAttention gains a window argument for sliding-window attention.

  • With window > 0, the token at position p attends only to positions p-window+1 through p. This is the masking used by OLMo 3's sliding_attention layers (window 4096 on three of every four layers) and by Mistral 7B v0.1 (window on every layer). It matches transformers' sliding-window causal mask.
  • window == 0 keeps today's unlimited causal attention.

Every existing caller (mistral, mistral4, nemotronh) passes 0, so behavior is unchanged until a model asks for a window; config.Model.SlidingWindowFor (#115) supplies the per-layer value.

What changes

  • Interface: PagedAttention(dst, q, kCache, vCache, blockTables, seqLens, queryLens, scale, window).
  • CPU reference: each query token's visible range starts at max(0, end-window) instead of 0.
  • CUDA:
    • Both kernels, paged_attention_kernel and paged_attention_flash_kernel, start their position loop at the same point. The flash kernel's warps start from there; a warp left with no positions contributes nothing to the combine, as it already does for short spans.
    • The launcher's flash dispatch compares the longest windowed span with the threshold, not the longest context.
  • Sizing backend: signature only.

The KV cache still stores every position. Freeing blocks a window can no longer reach would cut memory for long sequences; that is a separate change.

Verification

  • go build ./..., go vet ./..., go test ./... pass; gofmt is clean.
  • TestPagedAttentionVsReference (CPU) now runs at windows 0, 2, 5 and the full context against a float64 reference over the windowed positions, on scattered out-of-order blocks with GQA. Window = context matches unlimited.
  • New TestCUDAPagedAttentionWindow compares the GPU against the CPU reference:
    • batch: a 5-token prefill chunk over a 300-token context plus a decode row over 40;
    • windows: 16, 100 and 4096;
    • kernels: headDim 64 (flash kernel) and 66 (original kernel).
  • CUDA tests run on both GPUs of the development box:
package RTX PRO 6000 Blackwell (sm_120) RTX 3070 (sm_86)
internal/backend/cuda 45 pass 44 pass, 1 skip (TestCUDAMatMulNVFP4W8A8: no FP8 GEMM on sm_86)
internal/model/mistral 22 pass 22 pass
internal/model/nemotronh 13 pass 13 pass
internal/model/mistral4 6 pass 6 pass

No throughput was measured. With window == 0 the kernels' loop starts at 0, as before.

🤖 Generated with Claude Code

`backend.Ops.PagedAttention` gains a `window` argument for sliding-window attention. - With `window > 0`, the token at position p attends only to positions p-window+1 through p. This is the masking used by OLMo 3's `sliding_attention` layers (window 4096 on three of every four layers) and by Mistral 7B v0.1 (window on every layer). It matches transformers' sliding-window causal mask. - `window == 0` keeps today's unlimited causal attention. Every existing caller (`mistral`, `mistral4`, `nemotronh`) passes 0, so behavior is unchanged until a model asks for a window; `config.Model.SlidingWindowFor` (#115) supplies the per-layer value. ## What changes - **Interface:** `PagedAttention(dst, q, kCache, vCache, blockTables, seqLens, queryLens, scale, window)`. - **CPU reference:** each query token's visible range starts at `max(0, end-window)` instead of 0. - **CUDA:** - Both kernels, `paged_attention_kernel` and `paged_attention_flash_kernel`, start their position loop at the same point. The flash kernel's warps start from there; a warp left with no positions contributes nothing to the combine, as it already does for short spans. - The launcher's flash dispatch compares the longest windowed span with the threshold, not the longest context. - **Sizing backend:** signature only. The KV cache still stores every position. Freeing blocks a window can no longer reach would cut memory for long sequences; that is a separate change. ## Verification - `go build ./...`, `go vet ./...`, `go test ./...` pass; gofmt is clean. - `TestPagedAttentionVsReference` (CPU) now runs at windows 0, 2, 5 and the full context against a float64 reference over the windowed positions, on scattered out-of-order blocks with GQA. Window = context matches unlimited. - New `TestCUDAPagedAttentionWindow` compares the GPU against the CPU reference: - batch: a 5-token prefill chunk over a 300-token context plus a decode row over 40; - windows: 16, 100 and 4096; - kernels: headDim 64 (flash kernel) and 66 (original kernel). - CUDA tests run on both GPUs of the development box: | package | RTX PRO 6000 Blackwell (sm_120) | RTX 3070 (sm_86) | | --- | --- | --- | | `internal/backend/cuda` | 45 pass | 44 pass, 1 skip (`TestCUDAMatMulNVFP4W8A8`: no FP8 GEMM on sm_86) | | `internal/model/mistral` | 22 pass | 22 pass | | `internal/model/nemotronh` | 13 pass | 13 pass | | `internal/model/mistral4` | 6 pass | 6 pass | No throughput was measured. With `window == 0` the kernels' loop starts at 0, as before. 🤖 Generated with [Claude Code](https://claude.com/claude-code)
feat(backend): sliding-window PagedAttention
All checks were successful
ci / test_and_build (pull_request) Successful in 51s
f2a07123aa
PagedAttention gains a window argument: when positive, the token at
position p attends to positions p-window+1 through p only, the masking
OLMo 3's sliding_attention layers (and Mistral 7B v0.1's) use; 0 keeps
today's unlimited causal attention. The CPU reference starts each
token's visible range at the window, and both CUDA kernels start their
position loop there; the launcher's flash dispatch measures the longest
windowed span rather than the longest context.

Every caller passes 0, so behavior is unchanged until a model asks for
a window. The KV cache still keeps every position; evicting the ones a
window no longer reaches is a separate memory optimization.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Collaborator

Automated review by pr-reviewer v0.54.0 | Safety Check | Nemotron 3 Nano | tracking id r-c1e93c-f4307c
This is an AI-generated review and may contain mistakes.

Status: ❌ Failed


Review failed. Tracking id r-c1e93c-f4307c — see logs for details.

Comment @pr-reviewer-bot retry to try again.

<!-- pr-reviewer:review --> *Automated review by [pr-reviewer](https://git.brooktrails.org/brooktrails/pr-reviewer) v0.54.0 | Safety Check | Nemotron 3 Nano | tracking id `r-c1e93c-f4307c`* *This is an AI-generated review and may contain mistakes.* **Status:** ❌ Failed --- Review failed. Tracking id `r-c1e93c-f4307c` — see logs for details. Comment `@pr-reviewer-bot retry` to try again.
rcsheets deleted branch feat/sliding-window-attention 2026-10-04 06:23:27 +00:00
Sign in to join this conversation.
No reviewers
No labels
No milestone
No project
No assignees
2 participants
Notifications
Due date
The due date is invalid or out of range. Please use the format "yyyy-mm-dd".

No due date set.

Dependencies

No dependencies set

Reference
brooktrails/gllm!116
No description provided.