3Day 3

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.

Not measuredEXP-018

A sliding window throws distant tokens away completely. Compression keeps a lossy summary of them instead. Every block of bb old keys becomes one averaged key, and the same for values. The cache for old positions shrinks by a factor of bb, and a query can still reach far back, blurrily.

Definition

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.

Playground
Block size
columns 0–3: pooled blocks, then 16 exact keys16 × 20
b0b1b2b30123456789101112131415q0q1q2q3q4q5q6q7q8q9q10q11q12q13q14q15
Pooled entries4mean of 4 keys each
Cache at 4K, GQA-24.00 MiBevery position cached
Cache at 4K, window 256, block 41.19 MiB256 exact + pooled blocks

A query reads its most recent 4 keys exactly. Older keys reach it only as block averages, and only once the whole block has left the window. The cache figures use the Day 2 shape through kv_cache_bytes.

16 positions. The first columns, b0 onward, are pooled blocks. The rest are the exact keys under a sliding window. Filled cells are allowed. A test checks this mask against octlm's compressed_mask. The cache figures come from kv_cache_bytes at the Day 2 shape.

What happened?

octlm/model.pyline 194
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)
Playground
keys, one bar per positionwhat the last query reads
Columns the last query reads10full causal: 24
Pooled blocks4orange, one mean each
Positions lost2a partial block is dropped

Old keys are averaged in blocks, so the query sees their rough shape but not any single position. Recent keys, in blue, stay exact. A copy task that needs one old token exactly is the case pooling can blur, which is what the needle probe was built to test. The cache holds window + ⌊(T − window) / block⌋ positions instead of T.

Twenty-four made-up key values. Old keys are averaged per block, recent ones stay exact.
A problem with RoPE

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.

octlm/model.pyline 279
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)
The probe

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:

  1. Take a block of real text and plant a random 8-token sequence at depth pp.
  2. Repeat the same 8 tokens at the end of the block.
  3. Measure the model's loss on that final copy.
  4. Compare it with a control block where depth pp 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.

octlm/day3.pyline 104
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 results

The note is explicit that this is not the long-context needle test the literature runs.

The hypothesis

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.

Status

Built, tested, not scheduled

The plan changed before the run. The code stays behind default-off flags.

Skipped

What we did not build, and why