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
|
Block-diagonal causal mask: q attends to k iff same document and k <= q. |
|
Return the |
- kempnerforge.model.masking.flex_attention_fn(compiled)[source]¶
Return the
flex_attentioncallable 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:
- 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=Nonebroadcasts the mask over heads, which is what lets grouped-query attention passenable_gqa=Trueinstead 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_lenmust be at leastFLEX_BLOCK_SIZE: below it an Inductor-compiled model silently leaks attention across document boundaries, soJobConfig.validaterejects that configuration outright. This kernel is not the culprit – it is exact at those lengths, as is the same model underbackend="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
BlockMaskover(batch, seq_len, seq_len), head-broadcast.- Return type: