3Day 3

Sparse attention

A sliding window lets each token read only its recent neighbors, and strided columns add a few long-range links. The masks are built and tested, and the run was dropped when the plan changed.

Not measuredEXP-017

Full causal attention compares every token with every earlier token. The work and the score memory both grow with T2T^2, and in a trained model most of those scores end up near zero. Sparse attention decides in advance which scores to skip. The question for EXP-017 was what that costs on code, which refers further back than prose: a variable defined 200 lines up is still in scope.

Definition

A window, plus a few global columns

A sliding window of size ww lets query ii read keys jj with 0≤i−j<w0 \le i - j < w. Each row has at most ww allowed scores, so the cost grows with T⋅wT \cdot w instead of T2T^2. Longformer and Mistral use this pattern.

A window loses everything older than ww. A stride ss adds some of it back cheaply. Every key whose position is a multiple of ss stays visible to all later queries. Those columns act as landmarks that any token can reach.

allowed(i,j)=(j≤i)∧(i−j<w  ∨  j mod s=0)\text{allowed}(i, j) = (j \le i) \wedge \big( i - j < w \;\vee\; j \bmod s = 0 \big)
Playground
attention_mask(T, window, stride)16 × 16
0123456789101112131415q0q1q2q3q4q5q6q7q8q9q10q11q12q13q14q15
Density42.6%kept share of the causal triangle
Scores per row, last position4full causal: 16
Oldest key the last row reads12

A window keeps the 4 most recent keys per query. With a stride, every column divisible by it stays visible to all later rows, a cheap long-range path. Any mask other than plain causal makes PyTorch's SDPA leave the flash kernel, which Day 3 would have measured.

The mask octlm builds, row by row. Filled cells are allowed scores, hatched cells are skipped. Density is the share of the full causal triangle that remains. A test checks this mask and its density against octlm's attention_mask and mask_density.

What happened?

In octlm

One mask function for every pattern

The Day 1 causal mask and the Day 3 sparse masks come from the same function. With window = 0 it returns the plain triangle:

octlm/model.pyline 178
def attention_mask(
    length: int, window: int, stride: int, device: torch.device | None = None
) -> Tensor:
    """Causal mask, optionally narrowed to a sliding window with strided global columns."""
    rows = torch.arange(length, device=device).unsqueeze(1)
    columns = torch.arange(length, device=device).unsqueeze(0)
    distance = rows - columns
    allowed = distance >= 0
    if window:
        local = distance < window
        if stride:
            local = local | (columns % stride == 0)
        allowed = allowed & local
    return allowed.view(1, 1, length, length)

The model passes that mask as a boolean attn_mask into scaled_dot_product_attention. That is correct on every backend and every device, and it has a cost the note wrote down as a measurement rather than a defect. The moment an explicit mask replaces is_causal=True, SDPA can no longer use the flash kernel on CUDA and falls back to a slower backend. PyTorch's flex_attention can compile a block-sparse mask into a fast kernel, but it needs torch.compile, and chasing it was not what EXP-017 was for.

A small probe, flash_accepts, records which backend actually ran. It must run on the training device, because PyTorch's CPU flash path accepts an explicit mask and the CUDA kernel does not, so a probe on the laptop would report the wrong backend for a GPU run.

Playground
full causalwindow
1.6e+46.6e+42.6e+51.0e+64.2e+61.7e+71282565121,0242,0484,0968,192sequence lengthscores per headfull causalwindow
Hover the chart to read values.
Scores at 4,0961.02e+6full: 8.39e+6
Density12.1%

On log axes, full causal attention climbs with slope 2. A window climbs with slope 1 once the sequence is longer than the window: each query reads at most w keys. Global stride columns add back a quadratic term, T²/(2s), so a small stride gives up most of the saving at long lengths. Below the window length the two lines are the same, so a window only pays at lengths well past it.

Scores computed per head as the sequence grows, exact counts from the mask rule above. Both axes are logarithmic, base 2.
The hypothesis

Worse on code, and no faster at this length

The note's hypothesis had two parts:

  1. At a 512-token trained context, a 128-token window loses measurable bits per byte on code and less on prose.
  2. The windowed run is no faster than full causal. Losing the flash kernel costs more than the skipped scores save at 512 tokens. The speed win belongs to lengths this lab had not trained at.

Baseline: full causal attention at the same budget and context. Stop condition: windows of 128 and 256, strides of 0 and 64, three seeds, reporting code and prose bits per byte, seconds per step, the backend that ran and the computed density.

Status

Built, tested, not scheduled

The plan changed before any Day 3 run. None of the sparse patterns serve the tool-use harness the new plan builds toward. The code stays behind flags that default to off, so every other model is unaffected, and the tests for the masks still run.

Skipped

What we did not build, and why