Compressed attention and the copy probe
Keep the recent keys exactly and replace older ones with block averages, for a fraction of the cache. Day 3 built it, fixed a subtle problem with RoPE, and designed a probe that a 3M model could actually pass.
A sliding window throws distant tokens away completely. Compression keeps a lossy summary of them instead. Every block of old keys becomes one averaged key, and the same for values. The cache for old positions shrinks by a factor of , and a query can still reach far back, blurrily.
Two kinds of columns, pooled blocks and exact recent keys
DeepSeek's Native Sparse Attention runs three branches and mixes them with a learned gate: a compressed branch that mean-pools consecutive key-value blocks, a selected branch that picks the top-k blocks by their compressed scores, and a sliding window of exact recent tokens.
EXP-018 built two of the three: the compressed branch and the window. Each query attends over the pooled blocks first, then the exact keys. A pooled block becomes visible only once all of its tokens have left the query's window, so no token is ever counted twice.
What happened?
- Early rows see only exact keys. Their past fits inside the window.
- Later rows gain block columns one block at a time, as each block clears the window.
- A larger block means fewer pooled columns and a smaller cache, and each column averages more tokens, so it says less about any one of them.
- At 4K positions with a 256 window, block 4 cuts the Day 2 cache from 4 MiB to about 1.2 MiB.
def compressed_mask(length: int, window: int, block: int, device: torch.device | None) -> Tensor:
"""Columns are [pooled blocks, exact tokens]. A block is visible once it clears the window."""
rows = torch.arange(length, device=device).unsqueeze(1)
blocks = torch.arange(length // block, device=device).unsqueeze(0)
block_visible = (blocks + 1) * block <= rows - window + 1
exact = attention_mask(length, window, 0, device)[0, 0]
return torch.cat((block_visible, exact), dim=1).view(1, 1, length, -1)Pool before rotating, not after
With RoPE, each key has already been rotated by its position. Averaging rotated keys averages vectors pointing in different directions. The average is shorter than its parts, by an amount that depends on how spread out the block's angles are, and that biases the scores by block position.
The note's fix is to pool the keys before rotation, then rotate each pooled block once, at the position of its midpoint. NSA solves the same problem with an encoding inside each block. The note records the midpoint rotation as a cheaper version of that idea and a deviation from the paper.
def _compress(
self, query: Tensor, key: Tensor, value: Tensor, rope: tuple[Tensor, Tensor] | None
) -> tuple[Tensor, Tensor, Tensor]:
"""Pool before rotating, then rotate each pooled block at its midpoint position."""
block = self.config.kv_compress_block
pooled_key, pooled_value = pool_blocks(key, block), pool_blocks(value, block)
if rope is not None:
cosine, sine = rope
midpoints = torch.arange(pooled_key.shape[2], device=key.device) * block + block // 2
midpoints = midpoints.clamp(max=cosine.shape[0] - 1)
pooled_key = apply_rope(pooled_key, cosine[midpoints], sine[midpoints])
query, key = apply_rope(query, cosine, sine), apply_rope(key, cosine, sine)
return query, torch.cat((pooled_key, key), 2), torch.cat((pooled_value, value), 2)A needle test a 3M model can pass
PLAN.md asked for "retrieval accuracy on a needle style synthetic test". The usual needle test hides a fact in a long document and asks a question about it. A 3.3 million parameter base model cannot answer questions, so that test would measure only its inability to follow an instruction. The note called the substitute a judgment call and flagged it everywhere it appears.
The substitute is a copy probe:
- Take a block of real text and plant a random 8-token sequence at depth .
- Repeat the same 8 tokens at the end of the block.
- Measure the model's loss on that final copy.
- Compare it with a control block where depth holds different random tokens, so the final 8 tokens appear for the first time.
A model that carried the earlier sequence forward predicts the repeat far better than a first sighting. That skill, copying a span it has seen, is called induction, and it appears even in tiny models. It disappears when the mechanism that carries the reference is cut. So the probe should show full attention recovering the needle at every depth, the window recovering it only within reach, and compression recovering it partially everywhere.
def copy_probe(
model: Decoder, blocks: TokenBlocks, vocab_size: int, depths: list[int]
) -> list[dict[str, object]]:
"""Plant a random span at `depth`, repeat it at the end, and read the NLL on the repeat.
A prompt-and-answer needle test measures instruction following, which a base model at this
scale does not have. Recalling a span it has already seen is the part it can do.
"""
generator = torch.Generator().manual_seed(PROBE_SEED)
filler = blocks.inputs[0].clone()
length = filler.shape[0]
results = []
for depth in depths:
if depth + NEEDLE_LENGTH > length - NEEDLE_LENGTH:
continue
needle = torch.randint(256, vocab_size, (NEEDLE_LENGTH,), generator=generator)
decoy = torch.randint(256, vocab_size, (NEEDLE_LENGTH,), generator=generator)
start = length - NEEDLE_LENGTH
planted, control = filler.clone(), filler.clone()
planted[depth : depth + NEEDLE_LENGTH] = needle
control[depth : depth + NEEDLE_LENGTH] = decoy
planted[start:] = needle
control[start:] = needle
# `_span_nll` reads targets, which are the inputs shifted by one.
span = (start - 1, length - 1)
recalled = _span_nll(model, planted, *span)
unseen = _span_nll(model, control, *span)
results.append(
{
"depth": depth,
"gain": unseen - recalled,
"needle_nll_recalled": recalled,
"needle_nll_unseen": unseen,
"type": "copy_probe",
}
)
return resultsThe note is explicit that this is not the long-context needle test the literature runs.
What the run would have tested
Block compression cuts the cache by roughly the block size outside the window, and costs less on code than a plain window with the same cache budget. On the copy probe, full attention recovers the needle at every depth, the window only inside itself, and compression partially at every depth.
Stop condition: blocks of 4 and 8 against an uncompressed control, three seeds, with the probe swept across depth for all three configurations. The probe would reuse the first seed's trained models rather than train its own, which saved about 12 minutes of T4 time for no loss of information.
Built, tested, not scheduled
The plan changed before the run. The code stays behind default-off flags.
What we did not build, and why
- NSA's selection branch and gate. Selection is where the paper's hardware-aligned kernel work lives. A top-k over blocks at 3.3M parameters would have measured our implementation of top-k, not the idea.
- A question-answer needle test. Replaced by the copy probe, for the reason above.