kempnerforge.model.masking

BlockMask builders for FlexAttention-based self-attention.

FlexAttention replaces the dense (B, 1, S, S) boolean mask that the packed path otherwise hands to F.scaled_dot_product_attention: the mask predicate is compiled into the attention kernel and fully-masked blocks are skipped rather than computed. For document packing the mask is block-diagonal and mostly zeros, so that is the difference between paying for the full S x S attention and paying only for the blocks inside a document.

Functions

build_doc_causal_block_mask(doc_ids, device)

Block-diagonal causal mask: q attends to k iff same document and k <= q.

flex_attention_fn(compiled)

Return the flex_attention callable to use.

kempnerforge.model.masking.flex_attention_fn(compiled)[source]

Return the flex_attention callable to use.

Parameters:

compiled (bool) – Whether to return the torch.compile-wrapped kernel. True on CUDA; False on CPU, where the eager decomposition keeps unit tests off Inductor’s C++ codegen path. Note torch 2.11 has no CPU backward for FlexAttention, so the CPU path is forward-only.

Return type:

Callable[[…], Any]

kempnerforge.model.masking.build_doc_causal_block_mask(doc_ids, device)[source]

Block-diagonal causal mask: q attends to k iff same document and k <= q.

Built once per forward and shared by every layer. H=None broadcasts the mask over heads, which is what lets grouped-query attention pass enable_gqa=True instead of materializing repeated K/V heads, and what keeps the mask correct under tensor parallelism, where each rank holds only a shard of the heads.

seq_len must be at least FLEX_BLOCK_SIZE: below it an Inductor-compiled model silently leaks attention across document boundaries, so JobConfig.validate rejects that configuration outright. This kernel is not the culprit – it is exact at those lengths, as is the same model under backend="eager" or "aot_eager"; only Inductor codegen diverges.

Parameters:
  • doc_ids (torch.Tensor) – Per-token document ids, shape (batch, seq_len).

  • device (torch.device) – Device on which to materialize the BlockMask.

Returns:

A BlockMask over (batch, seq_len, seq_len), head-broadcast.

Return type:

torch.nn.attention.flex_attention.BlockMask