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.
Full causal attention compares every token with every earlier token. The work and the score memory both grow with , 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.
A window, plus a few global columns
A sliding window of size lets query read keys with . Each row has at most allowed scores, so the cost grows with instead of . Longformer and Mistral use this pattern.
A window loses everything older than . A stride adds some of it back cheaply. Every key whose position is a multiple of stays visible to all later queries. Those columns act as landmarks that any token can reach.
What happened?
- With the window at 0 the mask is the full causal triangle, density 100 percent.
- A window of 4 turns it into a band. The last row reads only 4 keys, and density falls as the sequence grows, because the band's width is fixed while the triangle's area grows with .
- A stride adds vertical stripes. Every -th column stays visible all the way down.
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:
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.
Worse on code, and no faster at this length
The note's hypothesis had two parts:
- At a 512-token trained context, a 128-token window loses measurable bits per byte on code and less on prose.
- 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.
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.
What we did not build, and why
flex_attentionkernels. Needstorch.compile, and the Day 3 GPU is a T4, where the gain was untested.- Dilated or random patterns. The window and stride answered the question the plan asked.
- The read on how Longformer and Mistral choose a window size. Listed in the note as still to read, and not read before the plan changed.