Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

### Added

- **FlexAttention backend for packed sequences** (`model.attention_backend = "flex"`). Document packing built a dense `(B, 1, S, S)` boolean mask *in every layer* and handed it to `F.scaled_dot_product_attention`; an explicit `attn_mask` is not a FlashAttention-2 shape, so SDPA fell back to the mem-efficient/math kernel and packing gave up throughput to save padding. `"flex"` instead builds one FlexAttention `BlockMask` per forward, shared by every layer, so the mask predicate compiles into the attention kernel and fully-masked blocks are skipped rather than computed. The default `"sdpa"` is bit-identical to before, and the flag only does anything when `data.pack_sequences = true` — unpacked batches take the `is_causal` SDPA fast path under either backend.
- `kempnerforge/model/masking.py` (new): `build_doc_causal_block_mask` (block-diagonal causal; `H=None` so the mask broadcasts over heads, which is what keeps it correct under tensor parallelism's sharded head counts) and `flex_attention_fn`, which compiles `flex_attention` once per process on CUDA and stays eager on CPU.
- `kempnerforge/model/attention.py`: a `block_mask` branch ahead of the dense-mask path. GQA goes through `enable_gqa` instead of `repeat_interleave`, so the flex path never materializes the repeated K/V heads. `capture_attention_weights` raises under flex — the kernel fuses the mask and never forms an attention-weight matrix to capture.
- `kempnerforge/model/transformer.py`: builds the `BlockMask` once per forward and threads it to every block, replacing the per-layer dense mask. Raises when `doc_ids` is shorter than the sequence reaching attention (the VLM image-prefix case, where packing is unsupported).
- `kempnerforge/config/model.py`: `attention_backend`, validated to `{"sdpa", "flex"}` with `head_dim >= 16` required under flex (8 fails to lower with "NYI: embedding dimension") and a warning that `sdpa_backend` is ignored on that path. Deliberately **no upper bound**: head_dim 256/320/384/512 were all measured correct, forward and backward, to ~2e-6 against a dense reference. The usable set is not an interval either -- head_dim 192 fails to compile on H200 / torch 2.11 ("No valid triton configs", shared-memory exhaustion) while 128 and 256 on either side of it are fine -- so it is not expressed as a range; a config-time check cannot predict which sizes Triton can tile on a given GPU, and that failure is a loud compile-time error rather than silent corruption.
- `pyproject.toml`: torch floor `>=2.4` -> `>=2.6` (FlexAttention needs >=2.5; `enable_gqa` and `BlockMask` ergonomics want >=2.6). `uv.lock` already resolves 2.11.
- Tests: `tests/unit/test_masking.py` (mask vs dense reference, `seq_len` not a multiple of the 128 block size, forward parity, `enable_gqa` equivalence), `tests/unit/test_packing.py::TestPackedModelFlexBackend`, `tests/unit/test_config.py`, and `tests/integration/test_compile.py::TestFlexCompile` (GPU: gradient parity, cross-document isolation, and a recompile-count guard).
- `kempnerforge/config/job.py`: `attention_backend="flex"` requires `train.seq_len >= 128` (`FLEX_BLOCK_SIZE`). Below that a compiled model silently leaks attention across document boundaries while eager stays correct: measured as exact (0.0 drift) at 128/129/130/160/200/256/300/384/512/1000/2048 and leaking at 32/64/96/120/127, for head_dim 32, 64 and 128 alike -- a sequence-length effect, not a head-dimension one. Localized to **Inductor codegen**: the same model under `backend="eager"` or `backend="aot_eager"` is exact (0.0 drift) and only `backend="inductor"` leaks, so dynamo tracing, AOTAutograd and the kernel are all innocent -- bare compiled `flex_attention` also matches a dense reference at those lengths across head_dim 16/32/64/128. It additionally needs graph scale: `proj + rope + flex + o_proj` compiled alone is exact at seq_len 32, while a one-layer `Transformer` is not. **Not fixed by upgrading** -- reproduced with bit-identical drift on torch 2.11.0+cu128, 2.13.0+cu129 and 2.14.0+cu130 -- and the dense-mask `"sdpa"` path is **unaffected** at every length, eager and compiled alike, so existing packed runs are not at risk. Not reduced further; the bound is empirical. It costs nothing real either way -- a sub-block sequence has no block sparsity to exploit, and every shipped config uses `seq_len` 512, 1024, or 4096.
- **Known limitations**: FlexAttention has no CPU backward in torch 2.11 and raises as soon as an input requires grad, so `attention_backend="flex"` is CUDA-only in practice and CPU test coverage is forward-only under `torch.no_grad()`. `capture_attention_weights` is unsupported under flex, and `key_padding_mask` (VLM video) and MoT self-attention stay on the dense SDPA path.
- **Ragged grids in `attentional_pool`.** The `attentional_pool` connector now pools patch grids the window does not evenly divide, instead of rejecting them: the grid is padded to `ceil(grid/window) * window` and each partial edge window pools only its real patches, with the padded patches masked out of that window's K/V and of the mean-query. Reaches parity with `avgpool`, so any window pools any grid (e.g. a 3×3 window on a 14×14 SigLIP2 grid → 5×5 = 25 tokens, previously refused). `output_num_tokens` is `ceil(grid/window)²` for both pooling connectors; a divisible grid builds no mask and is bit-exact with before, so no existing config's token count or loss changes.
- `kempnerforge/model/adapter.py`: shared `_pad_grid_to_windows` helper; `AttentionalPoolAdapter.forward` masks padded patches out of each edge window; `AvgPoolAdapter.forward` masks its numerator through the same tensor.
- **Backward-incompatible signature change**: `pooled_token_count` loses its keyword-only `require_divisible` parameter and `DIVISIBLE_ONLY_POOL_TYPES` is removed. The parameter existed only to raise on ragged grids, which is the behaviour being removed, so no shim is provided; out-of-tree callers passing `require_divisible=` must drop the argument.
Expand Down
14 changes: 13 additions & 1 deletion codebase-map.json
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
"meta": {
"description": "KempnerForge codebase map — files, classes, functions, tests, and registry entries",
"updated": "2026-05-11",
"note": "Keep in sync as phases are implemented. Latest landings: VLM (Joint-Decoder, Cross-Attention, Mixture-of-Transformers) + freeze schedule + vision encoder registry + checkpoint helpers (peek_saved_step, flush_pending_save, cross-arch freeze-meta intersection) + FSDP2 helpers (_fsdp_wrap_transformer_blocks, _apply_fsdp_vlm, _build_vlm)."
"note": "Keep in sync as phases are implemented. Latest landings: VLM (Joint-Decoder, Cross-Attention, Mixture-of-Transformers) + freeze schedule + vision encoder registry + checkpoint helpers (peek_saved_step, flush_pending_save, cross-arch freeze-meta intersection) + FSDP2 helpers (_fsdp_wrap_transformer_blocks, _apply_fsdp_vlm, _build_vlm). Also: FlexAttention packed-attention backend (model/masking.py, ModelConfig.attention_backend)."
},
"source": {
"kempnerforge/config/registry.py": {
Expand Down Expand Up @@ -471,6 +471,18 @@
"tests/unit/test_hooks.py::TestAttentionWeightCapture"
]
},
"kempnerforge/model/masking.py": {
"functions": [
"flex_attention_fn(compiled) -> Callable # torch.compile'd on CUDA, eager on CPU",
"build_doc_causal_block_mask(doc_ids, device) -> BlockMask # block-diagonal causal"
],
"purpose": "FlexAttention BlockMask builders for packed sequences (model.attention_backend='flex'); replaces the dense (B, 1, S, S) SDPA mask",
"tested_by": [
"tests/unit/test_masking.py",
"tests/unit/test_packing.py::TestPackedModelFlexBackend",
"tests/integration/test_compile.py::TestFlexCompile"
]
},
"kempnerforge/model/mlp.py": {
"classes": {
"SwiGLUMLP": {
Expand Down
1 change: 1 addition & 0 deletions docs/configuration/config-sections.md
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ Architecture hyperparameters and MoE knobs.
| `init_std` | `float` | `0.02` | weight-init std (GPT-2 / Llama convention) |
| `model_type` | `str` | `"transformer"` | `model` registry key |
| `sdpa_backend` | `str` | `"auto"` | one of `"auto"`, `"flash"`, `"efficient"`, `"cudnn"`, `"math"` |
| `attention_backend` | `str` | `"sdpa"` | `"sdpa"` or `"flex"`; `"flex"` routes packed batches (`data.pack_sequences`) through FlexAttention's block-sparse mask. Requires `train.seq_len ≥ 128` and `head_dim ≥ 16` |

### MoE (all defaults produce a dense model)

Expand Down
19 changes: 19 additions & 0 deletions docs/configuration/validation-rules.md
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,25 @@ File:
- `dim % n_heads == 0` (head dim is integral).
- `n_heads % n_kv_heads == 0` (GQA replication factor is integral).
- `sdpa_backend ∈ {"auto", "flash", "efficient", "cudnn", "math"}`.
- `attention_backend ∈ {"sdpa", "flex"}`.
- `attention_backend = "flex"` requires `dim // n_heads ≥ 16` (below that the
Triton template fails to lower), and warns that `sdpa_backend` is ignored when
set. There is no upper bound — head_dim 256/320/384/512 are all verified
working. Note the usable set is not an interval: head_dim 192 fails to compile
on H200 / torch 2.11 while 128 and 256 are fine. That surfaces as a loud
compile-time error, not silent corruption, so it is documented rather than
guarded.
- `attention_backend = "flex"` requires `train.seq_len ≥ 128`. Below that an
Inductor-compiled model silently leaks attention across document boundaries.
Localized to Inductor codegen: the same model under `backend="eager"` or
`backend="aot_eager"` is exact, as is the FlexAttention kernel on its own, and
a single attention block compiled by Inductor is exact too — it takes a larger
graph to trigger. Reproduced identically on torch 2.11/2.13/2.14 (cu128 and
cu130), so it is not waiting on a release. The dense-mask `"sdpa"` path is
unaffected at every length, so existing packed runs are not at risk. Not
reduced further; the bound is empirical. It costs nothing, since a sequence
below the mask block size has no block sparsity to exploit.
- `pp > 1` rejects `data.pack_sequences` (pipeline stages never receive `doc_ids`).
- When `num_experts > 0` (MoE):
- `moe_top_k > 0`
- `moe_top_k ≤ num_experts`
Expand Down
10 changes: 9 additions & 1 deletion kempnerforge/config/job.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from kempnerforge.config.distributed import DistributedConfig
from kempnerforge.config.eval import EvalConfig
from kempnerforge.config.metrics import MetricsConfig
from kempnerforge.config.model import ModelConfig
from kempnerforge.config.model import FLEX_BLOCK_SIZE, ModelConfig
from kempnerforge.config.optimizer import OptimizerConfig
from kempnerforge.config.profiling import ProfilingConfig
from kempnerforge.config.scheduler import SchedulerConfig
Expand Down Expand Up @@ -131,6 +131,14 @@ def __post_init__(self) -> None:
"the VLM wrapper, so a [vlm] section (and [vision_encoder]) is required."
)

# Below one mask block, an Inductor-compiled model leaks attention across
# document boundaries; see FLEX_BLOCK_SIZE for the measurements.
if self.model.attention_backend == "flex" and self.train.seq_len < FLEX_BLOCK_SIZE:
raise ValueError(
f"attention_backend='flex' requires train.seq_len >= {FLEX_BLOCK_SIZE}, "
f"got {self.train.seq_len}. Use attention_backend='sdpa' for shorter sequences."
)

# Pipeline stages never receive doc_ids, so packing would silently train
# on cross-document context. See issue #204.
if self.distributed.pp > 1 and self.data.pack_sequences:
Expand Down
66 changes: 66 additions & 0 deletions kempnerforge/config/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,43 @@ class Activation(StrEnum):
relu = "relu"


# FlexAttention's mask block size, and the seq_len floor for
# attention_backend="flex". A compiled model leaks attention across document
# boundaries below it and is exact at or above it -- measured at seq_len
# 32/64/96/120/127 (leaking) against 128/129/130/160/200/256/300/384/512/1000/
# 2048 (exact), for head_dim 32, 64 and 128 alike, so this is a sequence-length
# effect rather than a head-dimension one.
#
# Localized to Inductor codegen: the same model compiled with backend="eager" or
# backend="aot_eager" is exact, and only backend="inductor" leaks -- so dynamo
# tracing, AOTAutograd and the FlexAttention kernel are all innocent (bare
# compiled flex_attention also matches a dense reference at these lengths). It
# further needs graph scale: proj + rope + flex + o_proj compiled on its own is
# exact at seq_len 32, while a one-layer Transformer is not. Not reduced further.
#
# Not fixed by upgrading: reproduced with bit-identical drift on torch 2.11.0+cu128,
# 2.13.0+cu129 and 2.14.0+cu130, so this bound is not a temporary workaround waiting
# on a release. The dense-mask SDPA path is unaffected at every length, eager and
# compiled alike -- this is specific to FlexAttention under Inductor, and existing
# packed runs on the default backend are not at risk.
#
# It costs nothing either way, since a sub-block sequence has no block sparsity
# to exploit.
FLEX_BLOCK_SIZE = 128

# Smallest head_dim FlexAttention's Triton template will lower; 8 fails with
# "NYI: embedding dimension". There is deliberately no upper bound: head_dim
# 256, 320, 384 and 512 were all measured correct (forward and backward, ~2e-6
# against a dense reference on H200 / torch 2.11). The set of usable sizes is
# not an interval, though -- head_dim 192 fails to compile there ("No valid
# triton configs", shared-memory exhaustion) while 128 and 256 on either side
# of it are fine. That failure is a loud compile-time error rather than silent
# corruption, so it is documented rather than guarded: a config-time check
# cannot predict which sizes Triton can tile on a given GPU, and rejecting an
# interval would block the many large head_dims that do work.
FLEX_MIN_HEAD_DIM = 16


@dataclass
class ModelConfig:
"""Architecture hyperparameters for a transformer model."""
Expand All @@ -41,6 +78,13 @@ class ModelConfig:
# SDPA backend: "auto" lets PyTorch select (recommended). Override to force
# a specific kernel for benchmarking or debugging.
sdpa_backend: str = "auto" # "auto", "flash", "efficient", "cudnn", "math"
# Attention backend. "sdpa" is the existing behavior: causal SDPA, or a dense
# (B, 1, S, S) mask when packing is on -- which drops SDPA off FlashAttention.
# "flex" routes packed batches through FlexAttention's sparse BlockMask, which
# skips fully-masked blocks instead of computing them. Only affects packed runs
# (data.pack_sequences); unpacked batches take the is_causal SDPA fast path
# under either backend.
attention_backend: str = "sdpa" # "sdpa", "flex"

# MoE (all defaults produce a dense model -- zero behavior change)
num_experts: int = 0 # 0 = dense, >0 = MoE
Expand Down Expand Up @@ -84,6 +128,28 @@ def __post_init__(self) -> None:
"Options: 'auto', 'flash', 'efficient', 'cudnn', 'math'"
)

# Attention backend validation
if self.attention_backend not in ("sdpa", "flex"):
raise ValueError(
f"Unknown attention_backend: '{self.attention_backend}'. Options: 'sdpa', 'flex'"
)
if self.attention_backend == "flex":
if self.head_dim < FLEX_MIN_HEAD_DIM:
raise ValueError(
f"attention_backend='flex' requires head_dim >= {FLEX_MIN_HEAD_DIM}, got "
f"{self.head_dim} (dim={self.dim} // n_heads={self.n_heads}). "
"FlexAttention's Triton template does not lower below that "
"(head_dim 8 fails with 'NYI: embedding dimension')."
)
if self.sdpa_backend != "auto":
import logging

logging.getLogger(__name__).warning(
"sdpa_backend=%r is ignored when attention_backend='flex' -- the flex "
"path does not route through the SDPA kernel selector.",
self.sdpa_backend,
)

# MoE validation
if self.num_experts > 0:
if self.moe_top_k <= 0:
Expand Down
42 changes: 40 additions & 2 deletions kempnerforge/model/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,15 +9,20 @@
from __future__ import annotations

import contextlib
from typing import TYPE_CHECKING

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel

from kempnerforge.model.masking import flex_attention_fn
from kempnerforge.model.norm import RMSNorm
from kempnerforge.model.position import apply_rope

if TYPE_CHECKING:
from torch.nn.attention.flex_attention import BlockMask

_SDPA_BACKENDS = {
"flash": SDPBackend.FLASH_ATTENTION,
"efficient": SDPBackend.EFFICIENT_ATTENTION,
Expand Down Expand Up @@ -122,6 +127,7 @@ def forward(
kv_cache: KVCache | None = None,
doc_ids: torch.Tensor | None = None,
key_padding_mask: torch.Tensor | None = None,
block_mask: BlockMask | None = None,
) -> torch.Tensor:
"""Forward pass.

Expand All @@ -138,10 +144,23 @@ def forward(
visual tokens of padded video frames). Combined with the causal
(and doc) mask; fully-masked query rows are unmasked to keep
softmax finite.
block_mask: Optional FlexAttention ``BlockMask`` encoding the same
block-diagonal causal mask ``doc_ids`` would build, but sparse:
fully-masked blocks are skipped by the kernel instead of
computed. Built once per forward by ``Transformer.forward``
when ``attention_backend="flex"``. Mutually exclusive with
``doc_ids``, ``key_padding_mask``, and ``kv_cache``.

Returns:
Output tensor of shape (batch, seq_len, dim).
"""
if block_mask is not None and self.capture_attention_weights:
raise NotImplementedError(
"capture_attention_weights requires attention_backend='sdpa'. "
"FlexAttention fuses the mask into the kernel and never materializes "
"the attention weight matrix."
)

batch, seq_len, _ = x.shape

# Project to Q, K, V
Expand Down Expand Up @@ -169,8 +188,10 @@ def forward(
if kv_cache is not None:
k, v = kv_cache.update(k, v)

# Expand KV heads for GQA: (batch, n_kv_heads, seq, dim) → (batch, n_heads, seq, dim)
if self.n_rep > 1:
# Expand KV heads for GQA: (batch, n_kv_heads, seq, dim) → (batch, n_heads, seq, dim).
# FlexAttention expands internally via enable_gqa, so the flex path skips
# materializing the repeated heads entirely.
if self.n_rep > 1 and block_mask is None:
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)

Expand All @@ -187,6 +208,23 @@ def forward(
q, k, v, seq_len, doc_ids, kv_cache, key_padding_mask
)
self.last_attention_weights = attn_weights.detach().cpu()
elif block_mask is not None:
# Packed sequences via FlexAttention. The BlockMask already encodes
# causal AND same-document, so doc_ids carries no extra information
# here. key_padding_mask is deliberately not folded in -- that path
# (VLM video) stays on the dense SDPA mask below.
assert kv_cache is None and key_padding_mask is None, (
"block_mask does not compose with kv_cache decode or key_padding_mask; "
"those cases use the dense SDPA mask."
)
out = flex_attention_fn(q.is_cuda)(
q,
k,
v,
block_mask=block_mask,
scale=self.head_dim**-0.5,
enable_gqa=self.n_rep > 1,
)
elif doc_ids is not None or key_padding_mask is not None:
# An explicit attn_mask is not a FlashAttention-2 shape, so SDPA falls
# back to the mem-efficient/math kernel here. The image-prefix video
Expand Down
Loading
Loading