diff --git a/CHANGELOG.md b/CHANGELOG.md index 6526f65..dffc3ab 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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. diff --git a/codebase-map.json b/codebase-map.json index a0f8433..04da538 100644 --- a/codebase-map.json +++ b/codebase-map.json @@ -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": { @@ -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": { diff --git a/docs/configuration/config-sections.md b/docs/configuration/config-sections.md index 9451a6e..b0d7a53 100644 --- a/docs/configuration/config-sections.md +++ b/docs/configuration/config-sections.md @@ -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) diff --git a/docs/configuration/validation-rules.md b/docs/configuration/validation-rules.md index 2700076..0b10791 100644 --- a/docs/configuration/validation-rules.md +++ b/docs/configuration/validation-rules.md @@ -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` diff --git a/kempnerforge/config/job.py b/kempnerforge/config/job.py index f87ac25..7527673 100644 --- a/kempnerforge/config/job.py +++ b/kempnerforge/config/job.py @@ -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 @@ -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: diff --git a/kempnerforge/config/model.py b/kempnerforge/config/model.py index 739c958..0321a9b 100644 --- a/kempnerforge/config/model.py +++ b/kempnerforge/config/model.py @@ -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.""" @@ -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 @@ -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: diff --git a/kempnerforge/model/attention.py b/kempnerforge/model/attention.py index b4b226c..da91fc3 100644 --- a/kempnerforge/model/attention.py +++ b/kempnerforge/model/attention.py @@ -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, @@ -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. @@ -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 @@ -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) @@ -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 diff --git a/kempnerforge/model/masking.py b/kempnerforge/model/masking.py new file mode 100644 index 0000000..0e9f00e --- /dev/null +++ b/kempnerforge/model/masking.py @@ -0,0 +1,95 @@ +"""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. +""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import Any + +import torch +from torch.nn.attention.flex_attention import BlockMask, create_block_mask, flex_attention + +# ``flex_attention`` only reaches its fused kernel under ``torch.compile``, and +# roughly half the shipped configs set ``train.compile_model = false``, so compile +# it here rather than relying on the outer model being compiled. Nesting this +# inside an outer ``torch.compile`` is fine -- Dynamo unwraps the inner context, +# and the two paths were measured to agree to ~4e-7 at every supported seq_len. +# +# Deliberately a plain module global rather than a cached factory: Dynamo warns +# that it ignores ``functools.lru_cache`` wrappers and traces the wrapped body +# directly, which it flags as a silent-incorrectness risk. +_FLEX_COMPILED = torch.compile(flex_attention, dynamic=False) +_CREATE_BLOCK_MASK_COMPILED = torch.compile(create_block_mask, dynamic=False) + + +def flex_attention_fn(compiled: bool) -> Callable[..., Any]: + """Return the ``flex_attention`` callable to use. + + Args: + compiled: 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 _FLEX_COMPILED if compiled else flex_attention + + +def build_doc_causal_block_mask(doc_ids: torch.Tensor, device: torch.device) -> BlockMask: + """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. + + Args: + doc_ids: Per-token document ids, shape ``(batch, seq_len)``. + device: Device on which to materialize the ``BlockMask``. + + Returns: + A ``BlockMask`` over ``(batch, seq_len, seq_len)``, head-broadcast. + """ + return _build_doc_causal_block_mask(doc_ids, device) # type: ignore[reportCallIssue] + + +@torch._dynamo.disable +def _build_doc_causal_block_mask(doc_ids: torch.Tensor, device: torch.device) -> BlockMask: + """``build_doc_causal_block_mask`` body, hidden from Dynamo. + + ``create_block_mask`` is not meant to be traced by an enclosing + ``torch.compile``. Disabling here costs one graph break at the top of the + model forward -- not one per layer, and not one per attention call. The + public wrapper above exists so callers (and pyright) see a real signature; + ``torch._dynamo.disable`` erases the one it wraps. + """ + batch, seq_len = doc_ids.shape + # int32 halves the index-load cost inside the mask kernel. The dataset emits + # int64; one sequence never holds anywhere near 2**31 documents. Moving to + # `device` here too, so the signature means what it says rather than quietly + # requiring the caller to have done it. + doc_ids = doc_ids.to(device=device, dtype=torch.int32) + + def mask_mod( + b: torch.Tensor, h: torch.Tensor, q_idx: torch.Tensor, kv_idx: torch.Tensor + ) -> torch.Tensor: + return (kv_idx <= q_idx) & (doc_ids[b, q_idx] == doc_ids[b, kv_idx]) + + # Compiled on CUDA: eager construction costs ~3 ms per forward regardless of + # batch, which is pure overhead at small per-rank work. CPU stays eager to + # keep unit tests off Inductor's C++ codegen path. + builder = _CREATE_BLOCK_MASK_COMPILED if device.type == "cuda" else create_block_mask + return builder(mask_mod, B=batch, H=None, Q_LEN=seq_len, KV_LEN=seq_len, device=device) diff --git a/kempnerforge/model/transformer.py b/kempnerforge/model/transformer.py index 20d26b8..32aa406 100644 --- a/kempnerforge/model/transformer.py +++ b/kempnerforge/model/transformer.py @@ -11,7 +11,7 @@ from __future__ import annotations -from typing import cast +from typing import TYPE_CHECKING, cast import torch import torch.nn as nn @@ -23,11 +23,15 @@ from kempnerforge.model.cross_attention import CrossAttentionBlock from kempnerforge.model.embedding import OutputHead, TokenEmbedding from kempnerforge.model.init import init_weights +from kempnerforge.model.masking import build_doc_causal_block_mask from kempnerforge.model.mlp import build_mlp from kempnerforge.model.modality import ModalityContext from kempnerforge.model.moe import MoEMLP, build_moe from kempnerforge.model.moma import ExpertChoiceMoE, MoMaBlock, MoMaFFN from kempnerforge.model.mot import MoTBlock + +if TYPE_CHECKING: + from torch.nn.attention.flex_attention import BlockMask from kempnerforge.model.norm import build_norm from kempnerforge.model.position import precompute_rope_frequencies @@ -88,6 +92,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: # Pre-norm attention with residual x = x + self.attention( @@ -97,6 +102,7 @@ def forward( kv_cache=kv_cache, doc_ids=doc_ids, key_padding_mask=key_padding_mask, + block_mask=block_mask, ) # Pre-norm MLP with residual x = x + self.mlp(self.mlp_norm(x)) @@ -424,6 +430,28 @@ def forward( cos = self._rope_cos[start_pos : start_pos + seq_len] # type: ignore[reportOptionalSubscript] sin = self._rope_sin[start_pos : start_pos + seq_len] # type: ignore[reportOptionalSubscript] + # Packed sequences under attention_backend="flex": build the sparse + # block-diagonal causal mask once here and share it across every layer, + # rather than the dense (B, 1, S, S) mask each Attention would otherwise + # rebuild. key_padding_mask (VLM video) is not folded into the BlockMask, + # so those batches stay on the dense SDPA path. + block_mask = None + if ( + doc_ids is not None + and key_padding_mask is None + and self.config.attention_backend == "flex" + ): + if doc_ids.shape[1] != seq_len: + raise ValueError( + f"doc_ids length ({doc_ids.shape[1]}) does not match the sequence " + f"length reaching attention ({seq_len}). Sequence packing is not " + "supported for arches that prepend image tokens to the residual stream." + ) + block_mask = build_doc_causal_block_mask(doc_ids, h.device) + # The BlockMask supersedes doc_ids; clearing it keeps the dense + # SDPA mask branch in Attention.forward unreachable on this path. + doc_ids = None + # MoMa path: single residual stream + shared SDPA + per-modality # MoE FFN groups. modality_ids tags every position and the # ``MoMaFFN`` uses these tags to dispatch tokens to per-modality @@ -495,7 +523,13 @@ def forward( for i, layer in enumerate(self.layers.values()): cache = kv_caches[i] if kv_caches is not None else None h = layer( - h, cos, sin, kv_cache=cache, doc_ids=doc_ids, key_padding_mask=key_padding_mask + h, + cos, + sin, + kv_cache=cache, + doc_ids=doc_ids, + key_padding_mask=key_padding_mask, + block_mask=block_mask, ) if ca_iter is not None and (i + 1) % self._ca_cadence == 0: ca = next(ca_iter, None) diff --git a/pyproject.toml b/pyproject.toml index 4a58905..1d6a5b2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ authors = [ ] requires-python = ">=3.12" dependencies = [ - "torch>=2.4", + "torch>=2.6", "wandb", "tensorboard", "transformers", diff --git a/tests/integration/test_compile.py b/tests/integration/test_compile.py index 699101e..58dec39 100644 --- a/tests/integration/test_compile.py +++ b/tests/integration/test_compile.py @@ -11,6 +11,7 @@ import torch import torch.nn.functional as F +from kempnerforge.config.model import FLEX_BLOCK_SIZE from kempnerforge.config.schema import ModelConfig from kempnerforge.model.transformer import Transformer @@ -87,3 +88,212 @@ def test_compiled_loss_matches_eager(self): ) assert abs(eager_loss.item() - compiled_loss.item()) < 1e-4 + + +@pytest.mark.gpu +class TestFlexCompile: + """FlexAttention packed-attention path, which only exists on CUDA. + + torch 2.11 has no CPU backward for flex and raises as soon as an input + requires grad, so every gradient and compiled-kernel assertion for that + backend lives here rather than in ``tests/unit/``. + + Every sequence here is at least ``FLEX_BLOCK_SIZE``. Below one block the + compiled kernel silently returns wrong results -- documents leak into each + other -- which is why ``JobConfig.validate`` rejects that configuration; + see ``tests/unit/test_config.py`` for the guard itself. + """ + + SEQ = 256 + + @staticmethod + def _layouts(seq: int) -> list[list[int]]: + half = seq // 2 + quarter = seq // 4 + return [ + [0] * seq, + [0] * half + [1] * (seq - half), + [0] * quarter + [1] * quarter + [2] * (seq - 2 * quarter), + [min(i * 8 // seq, 7) for i in range(seq)], + [0] * (seq - 1) + [1], + ] + + def _model(self, backend: str, seq: int | None = None, **overrides): + torch.manual_seed(0) + config = ModelConfig( + dim=128, + n_layers=2, + n_heads=4, + vocab_size=256, + max_seq_len=max(seq or self.SEQ, FLEX_BLOCK_SIZE), + attention_backend=backend, + **overrides, + ) + return Transformer(config).to("cuda") + + def _batch(self, layout_idx: int = 1, seq: int | None = None, seed: int = 0): + seq = seq or self.SEQ + gen = torch.Generator(device="cuda").manual_seed(seed) + tokens = torch.randint(0, 256, (2, seq), device="cuda", generator=gen) + row = self._layouts(seq)[layout_idx] + return tokens, torch.tensor([row, row], device="cuda") + + @staticmethod + def _perturb_first_doc(tokens: torch.Tensor, boundary: int) -> torch.Tensor: + perturbed = tokens.clone() + perturbed[:, :boundary] = (perturbed[:, :boundary] + 1) % 256 + return perturbed + + @pytest.mark.parametrize("overrides", [{}, {"n_kv_heads": 2}], ids=["mha", "gqa"]) + def test_forward_matches_sdpa(self, overrides): + """Flex reproduces the dense-mask SDPA logits it replaces.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + tokens, doc_ids = self._batch() + with torch.no_grad(): + out_sdpa = self._model("sdpa", **overrides).eval()(tokens, doc_ids=doc_ids) + out_flex = self._model("flex", **overrides).eval()(tokens, doc_ids=doc_ids) + torch.testing.assert_close(out_flex, out_sdpa, rtol=2e-5, atol=2e-5) + + @pytest.mark.parametrize("overrides", [{}, {"n_kv_heads": 2}], ids=["mha", "gqa"]) + def test_gradients_match_sdpa(self, overrides): + """Backward parity -- the assertion CPU cannot make (no flex CPU backward). + + Compared relative to each parameter's own gradient scale rather than with + a fixed atol: at S=256 gradients reach ~1.8e3, so an absolute bound is + meaningless. Measured against an fp64 CPU reference, flex is *closer* to + exact than the dense-mask SDPA path on every parameter (relative error + 3.1e-7..5.4e-7 vs SDPA's 3.2e-7..9.1e-7), and the two fp32 kernels differ + by 4.9e-7 of full scale. The 1e-5 bound below therefore carries ~20x + headroom while still catching real divergence, which shows up orders of + magnitude larger -- the pre-fix compiled bug moved logits by 4e-2. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + tokens, doc_ids = self._batch() + grads = {} + for backend in ("sdpa", "flex"): + model = self._model(backend, **overrides) + model(tokens, doc_ids=doc_ids).sum().backward() + grads[backend] = { + name: p.grad.detach().clone() + for name, p in model.named_parameters() + if p.grad is not None + } + + assert grads["sdpa"].keys() == grads["flex"].keys() + for name, reference in grads["sdpa"].items(): + actual = grads["flex"][name] + assert torch.isfinite(actual).all(), f"non-finite grad for {name}" + scale = reference.abs().max() + if scale == 0: + torch.testing.assert_close(actual, reference, rtol=0, atol=0) + continue + # Largest elementwise disagreement, as a fraction of full scale. + elementwise = ((actual - reference).abs().max() / scale).item() + # Whole-tensor agreement, which a single bad element cannot hide in. + relative_norm = ((actual - reference).norm() / reference.norm()).item() + assert elementwise < 1e-5, f"{name}: elementwise {elementwise:.2e} of scale {scale:.2e}" + assert relative_norm < 1e-5, f"{name}: relative norm {relative_norm:.2e}" + + @pytest.mark.parametrize("compiled", [False, True], ids=["eager", "compiled"]) + def test_cross_document_isolation(self, compiled): + """Perturbing document 0 must not move any document-1 logit, at all. + + The compiled case is a regression test: an earlier revision passed + eagerly and leaked under ``torch.compile``, which is the failure this + whole path exists to prevent and the one a loss curve would not reveal. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + model = self._model("flex").eval() + run = torch.compile(model) if compiled else model + tokens, doc_ids = self._batch(layout_idx=1) + boundary = self.SEQ // 2 + with torch.no_grad(): + base = run(tokens, doc_ids=doc_ids) + moved = run(self._perturb_first_doc(tokens, boundary), doc_ids=doc_ids) + torch.testing.assert_close(base[:, boundary:], moved[:, boundary:], rtol=0, atol=0) + assert not torch.allclose(base[:, :boundary], moved[:, :boundary]) + + @pytest.mark.parametrize("seq", [FLEX_BLOCK_SIZE, 200, 300]) + def test_sequence_lengths_around_the_block_size(self, seq): + """Exactly one block, and ragged lengths above it, all stay exact. + + ``create_block_mask`` rounds up to the block size, so 200 and 300 carry + a partial trailing block; the kernel must still isolate documents there. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + model = self._model("flex", seq=seq).eval() + compiled = torch.compile(model) + tokens, doc_ids = self._batch(layout_idx=1, seq=seq) + boundary = seq // 2 + with torch.no_grad(): + eager_out = model(tokens, doc_ids=doc_ids) + comp_out = compiled(tokens, doc_ids=doc_ids) + comp_moved = compiled(self._perturb_first_doc(tokens, boundary), doc_ids=doc_ids) + torch.testing.assert_close(comp_out, eager_out, rtol=2e-5, atol=2e-5) + torch.testing.assert_close(comp_out[:, boundary:], comp_moved[:, boundary:], rtol=0, atol=0) + + def test_no_recompilation_across_document_layouts(self): + """A fresh mask_mod closure per step must not retrigger compilation. + + ``build_doc_causal_block_mask`` defines ``mask_mod`` inline, closing over + that step's ``doc_ids``, so every step hands ``flex_attention`` a new + closure object. If Dynamo guarded on closure identity that would be a + recompile per step, which blows ``cache_size_limit`` and silently falls + back to *eager* flex -- the unfused decomposition that materializes the + full score matrix, i.e. slower and hungrier than the dense SDPA path this + replaces, with no error to notice. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + from torch._dynamo.utils import counters + + torch._dynamo.reset() + counters.clear() + + model = self._model("flex") + tokens, _ = self._batch() + layouts = self._layouts(self.SEQ) + + def step(layout_idx): + row = layouts[layout_idx] + doc_ids = torch.tensor([row, row], device="cuda") + model(tokens, doc_ids=doc_ids).sum().backward() + model.zero_grad(set_to_none=True) + + # Warm up on two different layouts: flex may compile one or two + # block-sparsity specializations up front, a fixed startup cost rather + # than a per-step recompile. + step(0) + step(1) + after_warmup = counters["stats"]["unique_graphs"] + assert after_warmup > 0, "flex_attention did not compile at all" + + for layout_idx in range(2, len(layouts)): + step(layout_idx) + + assert counters["stats"]["unique_graphs"] == after_warmup, ( + f"recompiled {counters['stats']['unique_graphs'] - after_warmup} time(s) across " + f"{len(layouts) - 2} document layouts -- the mask_mod closure is being guarded " + "on identity" + ) + + def test_compiled_model_matches_eager(self): + """flex composes with an outer torch.compile over the whole model.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + model = self._model("flex").eval() + tokens, doc_ids = self._batch() + with torch.no_grad(): + eager_out = model(tokens, doc_ids=doc_ids) + compiled_out = torch.compile(model)(tokens, doc_ids=doc_ids) + torch.testing.assert_close(compiled_out, eager_out, rtol=2e-5, atol=2e-5) diff --git a/tests/unit/test_config.py b/tests/unit/test_config.py index e67e351..3ae3b26 100644 --- a/tests/unit/test_config.py +++ b/tests/unit/test_config.py @@ -2,6 +2,7 @@ from __future__ import annotations +import logging import random import sys from pathlib import Path @@ -193,6 +194,43 @@ def test_sdpa_backend_rejects_unknown(self): with pytest.raises(ValueError, match="Unknown sdpa_backend"): ModelConfig(sdpa_backend="fa3") + # --- Attention backend config --- + + def test_attention_backend_default_is_sdpa(self): + assert ModelConfig().attention_backend == "sdpa" + + def test_attention_backend_accepts_valid_values(self): + for backend in ("sdpa", "flex"): + assert ModelConfig(attention_backend=backend).attention_backend == backend + + def test_attention_backend_rejects_unknown(self): + with pytest.raises(ValueError, match="Unknown attention_backend"): + ModelConfig(attention_backend="flash3") + + def test_flex_rejects_head_dim_below_16(self): + """FlexAttention's Triton template does not lower below head_dim 16.""" + with pytest.raises(ValueError, match="head_dim >= 16"): + ModelConfig(dim=64, n_heads=8, attention_backend="flex") + + def test_flex_accepts_head_dim_at_the_floor(self): + assert ModelConfig(dim=64, n_heads=4, attention_backend="flex").head_dim == 16 + + @pytest.mark.parametrize("head_dim", [256, 320, 384, 512]) + def test_flex_accepts_large_head_dim(self, head_dim): + """There is no upper bound: these were measured correct on H200. + + An earlier revision capped head_dim at 256, which rejected working + configurations. The usable set is not an interval either -- 192 fails to + compile while 128 and 256 are fine -- so it is deliberately not expressed + as a range; unsupported sizes surface as a loud compile-time error. + """ + config = ModelConfig(dim=head_dim, n_heads=1, attention_backend="flex") + assert config.head_dim == head_dim + + def test_sdpa_does_not_constrain_head_dim(self): + """The floor is a flex kernel limit, not an architecture limit.""" + assert ModelConfig(dim=64, n_heads=8).head_dim == 8 + # --------------------------------------------------------------------------- # TrainConfig @@ -506,6 +544,26 @@ def test_pp_without_packing_passes(self): config = JobConfig(distributed=DistributedConfig(pp=2, dp_shard=1)) config.validate(world_size=2) # Should not raise — PP alone is fine + def test_flex_rejects_seq_len_below_block_size_at_construction(self): + """Decidable from the config alone, so it fires in __post_init__.""" + with pytest.raises(ValueError, match="requires train.seq_len >= 128"): + JobConfig( + model=ModelConfig(attention_backend="flex"), + train=TrainConfig(seq_len=64), + ) + + def test_flex_accepts_seq_len_at_block_size(self): + config = JobConfig( + model=ModelConfig(attention_backend="flex"), + train=TrainConfig(seq_len=128), + ) + config.validate(world_size=1) # Should not raise — exactly one block is fine + + def test_sdpa_allows_short_seq_len(self): + """The bound is a flex kernel limit, not a general one.""" + config = JobConfig(train=TrainConfig(seq_len=64)) + config.validate(world_size=1) # Should not raise + def test_validate_vlm_seq_len_too_short(self): config = JobConfig( model=ModelConfig(max_seq_len=1024), @@ -554,6 +612,38 @@ def test_validate_vlm_plus_moe_supported(self): config.validate(world_size=1) # Should not raise. +class TestFlexSdpaBackendWarning: + """``sdpa_backend`` selects an SDPA kernel, which the flex path never uses, + so setting both is a silent no-op worth warning about. + + ``kempnerforge.metrics.logger._configure_root`` sets ``propagate = False`` + on the ``kempnerforge`` logger, which blocks pytest's caplog (attached at + the python root). Re-enable propagation around each test so caplog can see + records from ``kempnerforge.config.model``. Without this the tests pass in + isolation and fail once any earlier test has configured logging. + """ + + def setup_method(self): + self._kf_logger = logging.getLogger("kempnerforge") + self._old_propagate = self._kf_logger.propagate + self._kf_logger.propagate = True + + def teardown_method(self): + self._kf_logger.propagate = self._old_propagate + + def test_warns_when_sdpa_backend_set_under_flex(self, caplog): + with caplog.at_level(logging.WARNING, logger="kempnerforge.config.model"): + ModelConfig(attention_backend="flex", sdpa_backend="flash") + assert any( + "ignored when attention_backend='flex'" in r.getMessage() for r in caplog.records + ) + + def test_silent_when_sdpa_backend_set_under_sdpa(self, caplog): + with caplog.at_level(logging.WARNING, logger="kempnerforge.config.model"): + ModelConfig(attention_backend="sdpa", sdpa_backend="flash") + assert not any("ignored when attention_backend" in r.getMessage() for r in caplog.records) + + class TestHfEncoderOverrideWarning: """HF-backed encoders probe feature_dim / num_tokens from the loaded model. Non-zero TOML values override the probe at build time and diff --git a/tests/unit/test_masking.py b/tests/unit/test_masking.py new file mode 100644 index 0000000..29b2bd0 --- /dev/null +++ b/tests/unit/test_masking.py @@ -0,0 +1,133 @@ +"""Unit tests for FlexAttention BlockMask construction. + +Covers the mask predicate itself (against a dense reference) and forward +parity between the flex path and the dense-mask SDPA path it replaces. + +FlexAttention has no CPU backward in torch 2.11, so everything here is +forward-only; gradient parity is covered on GPU in +``tests/integration/test_compile.py``. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn.functional as F +from torch.nn.attention.flex_attention import create_mask + +from kempnerforge.model.masking import build_doc_causal_block_mask, flex_attention_fn + +CPU = torch.device("cpu") + + +def dense_doc_causal(doc_ids: torch.Tensor) -> torch.Tensor: + """Reference mask: causal AND same-document, as Attention builds it today.""" + seq_len = doc_ids.shape[1] + causal = torch.ones(seq_len, seq_len, dtype=torch.bool).tril() + return causal.unsqueeze(0) & (doc_ids.unsqueeze(2) == doc_ids.unsqueeze(1)) + + +def elementwise_mask(doc_ids: torch.Tensor) -> torch.Tensor: + """Expand the BlockMask's predicate to a dense (B, S, S) bool tensor.""" + batch, seq_len = doc_ids.shape + block_mask = build_doc_causal_block_mask(doc_ids, CPU) + return create_mask( + block_mask.mask_mod, B=batch, H=1, Q_LEN=seq_len, KV_LEN=seq_len, device=CPU + )[:, 0] + + +class TestBlockMaskStructure: + @pytest.mark.parametrize( + "doc_ids", + [ + pytest.param([[0, 0, 0, 0, 0, 0, 0, 0]], id="single_document"), + pytest.param([[0, 0, 0, 1, 1, 1, 1, 1]], id="two_documents"), + pytest.param([[0, 0, 1, 1, 1, 2, 2, 2]], id="three_documents"), + pytest.param([[0, 1, 2, 3, 4, 5, 6, 7]], id="every_token_its_own_document"), + pytest.param([[0, 0, 0, 1, 1, 1, 1, 1], [0, 1, 1, 1, 2, 2, 2, 2]], id="batched"), + ], + ) + def test_matches_dense_reference(self, doc_ids): + ids = torch.tensor(doc_ids) + assert torch.equal(elementwise_mask(ids), dense_doc_causal(ids)) + + def test_single_document_is_plain_causal(self): + ids = torch.zeros(1, 16, dtype=torch.long) + expected = torch.ones(16, 16, dtype=torch.bool).tril() + assert torch.equal(elementwise_mask(ids)[0], expected) + + def test_every_token_attends_to_itself(self): + """No query row is ever fully masked, so softmax cannot produce NaN.""" + ids = torch.tensor([[0, 1, 2, 3, 4, 5, 6, 7]]) + assert elementwise_mask(ids)[0].diagonal().all() + assert elementwise_mask(ids).any(dim=-1).all() + + @pytest.mark.parametrize("seq_len", [7, 63, 129, 200]) + def test_sequence_length_not_a_multiple_of_block_size(self, seq_len): + """BLOCK_SIZE is 128; ragged lengths must still mask exactly.""" + ids = torch.zeros(1, seq_len, dtype=torch.long) + ids[0, seq_len // 2 :] = 1 + assert torch.equal(elementwise_mask(ids), dense_doc_causal(ids)) + + def test_accepts_int64_doc_ids_from_the_dataset(self): + """_compute_packed_output emits int64; the builder casts to int32 internally.""" + ids = torch.tensor([[0, 0, 1, 1]], dtype=torch.int64) + assert build_doc_causal_block_mask(ids, CPU) is not None + assert torch.equal(elementwise_mask(ids), dense_doc_causal(ids)) + + +class TestFlexForwardParity: + """Flex output vs the dense-mask SDPA it replaces, in fp64 for a tight bound.""" + + def _qkv(self, batch, n_heads, seq_len, head_dim, seed=0): + gen = torch.Generator().manual_seed(seed) + shape = (batch, n_heads, seq_len, head_dim) + return tuple(torch.randn(shape, dtype=torch.float64, generator=gen) for _ in range(3)) + + def test_matches_dense_mask_sdpa(self): + doc_ids = torch.tensor([[0, 0, 0, 1, 1, 2, 2, 2]]) + q, k, v = self._qkv(1, 4, 8, 16) + block_mask = build_doc_causal_block_mask(doc_ids, CPU) + out_flex = flex_attention_fn(False)(q, k, v, block_mask=block_mask, scale=16**-0.5) + out_sdpa = F.scaled_dot_product_attention( + q, k, v, attn_mask=dense_doc_causal(doc_ids).unsqueeze(1) + ) + torch.testing.assert_close(out_flex, out_sdpa, rtol=1e-12, atol=1e-12) + + def test_single_document_matches_is_causal_sdpa(self): + doc_ids = torch.zeros(1, 8, dtype=torch.long) + q, k, v = self._qkv(1, 4, 8, 16, seed=1) + block_mask = build_doc_causal_block_mask(doc_ids, CPU) + out_flex = flex_attention_fn(False)(q, k, v, block_mask=block_mask, scale=16**-0.5) + out_sdpa = F.scaled_dot_product_attention(q, k, v, is_causal=True) + torch.testing.assert_close(out_flex, out_sdpa, rtol=1e-12, atol=1e-12) + + def test_enable_gqa_matches_manual_kv_expansion(self): + """Licenses skipping repeat_interleave on the flex path (attention.py).""" + doc_ids = torch.tensor([[0, 0, 0, 1, 1, 2, 2, 2]]) + gen = torch.Generator().manual_seed(2) + q = torch.randn(1, 4, 8, 16, dtype=torch.float64, generator=gen) + k = torch.randn(1, 2, 8, 16, dtype=torch.float64, generator=gen) + v = torch.randn(1, 2, 8, 16, dtype=torch.float64, generator=gen) + block_mask = build_doc_causal_block_mask(doc_ids, CPU) + flex = flex_attention_fn(False) + out_gqa = flex(q, k, v, block_mask=block_mask, scale=16**-0.5, enable_gqa=True) + out_expanded = flex( + q, + k.repeat_interleave(2, dim=1), + v.repeat_interleave(2, dim=1), + block_mask=block_mask, + scale=16**-0.5, + ) + torch.testing.assert_close(out_gqa, out_expanded, rtol=0, atol=0) + + +class TestFlexAttentionFn: + def test_returns_the_same_callable_for_repeated_calls(self): + """Cached so the compiled artifact is built once, not once per forward.""" + assert flex_attention_fn(False) is flex_attention_fn(False) + + def test_cpu_path_is_uncompiled(self): + from torch.nn.attention.flex_attention import flex_attention + + assert flex_attention_fn(False) is flex_attention diff --git a/tests/unit/test_packing.py b/tests/unit/test_packing.py index 657d7fb..5e05e0e 100644 --- a/tests/unit/test_packing.py +++ b/tests/unit/test_packing.py @@ -366,6 +366,83 @@ def test_gradient_flows_with_doc_ids(self, tiny_model): assert torch.isfinite(p.grad).all() +class TestPackedModelFlexBackend: + """Model-level parity for ``attention_backend="flex"``. + + FlexAttention raises on CPU as soon as an input requires grad -- even for a + forward, because autograd would need a backward node it does not have on CPU + in torch 2.11. Every case here therefore runs under ``torch.no_grad()``; + gradient and compiled-kernel coverage lives in + ``tests/integration/test_compile.py::TestFlexCompile``. + """ + + def _model(self, backend: str, **overrides): + from kempnerforge.config.schema import ModelConfig + from kempnerforge.model.transformer import Transformer + + torch.manual_seed(0) + config = ModelConfig( + dim=64, + n_layers=2, + n_heads=4, + vocab_size=256, + max_seq_len=32, + attention_backend=backend, + **overrides, + ) + return Transformer(config).eval() + + @pytest.mark.parametrize("overrides", [{}, {"n_kv_heads": 2}], ids=["mha", "gqa"]) + def test_matches_sdpa_backend(self, overrides): + """Same weights, same batch: flex reproduces the dense-mask SDPA output. + + The GQA case matters on its own -- flex passes ``enable_gqa`` instead of + materializing repeated K/V heads, so it exercises a different kernel path. + """ + tokens = torch.randint(0, 256, (2, 16), generator=torch.Generator().manual_seed(7)) + doc_ids = torch.tensor([[0] * 6 + [1] * 10, [0] * 16]) + with torch.no_grad(): + out_sdpa = self._model("sdpa", **overrides)(tokens, doc_ids=doc_ids) + out_flex = self._model("flex", **overrides)(tokens, doc_ids=doc_ids) + torch.testing.assert_close(out_flex, out_sdpa) + + def test_single_doc_matches_causal(self): + model = self._model("flex") + tokens = torch.randint(0, 256, (1, 8), generator=torch.Generator().manual_seed(1)) + with torch.no_grad(): + out_causal = model(tokens) + out_packed = model(tokens, doc_ids=torch.zeros(1, 8, dtype=torch.long)) + torch.testing.assert_close(out_causal, out_packed) + + def test_cross_doc_isolation(self): + """Perturbing document 0 must not move any document-1 output at all.""" + model = self._model("flex") + tokens = torch.randint(0, 256, (1, 16), generator=torch.Generator().manual_seed(3)) + doc_ids = torch.tensor([[0] * 6 + [1] * 10]) + perturbed = tokens.clone() + perturbed[0, :6] = (perturbed[0, :6] + 1) % 256 + with torch.no_grad(): + base = model(tokens, doc_ids=doc_ids) + moved = model(perturbed, doc_ids=doc_ids) + torch.testing.assert_close(base[:, 6:], moved[:, 6:], rtol=0, atol=0) + assert not torch.allclose(base[:, :6], moved[:, :6]) + + def test_no_doc_ids_is_bit_identical_to_sdpa(self): + """Unpacked batches build no BlockMask and take the is_causal SDPA path.""" + tokens = torch.randint(0, 256, (2, 16), generator=torch.Generator().manual_seed(5)) + with torch.no_grad(): + out_sdpa = self._model("sdpa")(tokens) + out_flex = self._model("flex")(tokens) + torch.testing.assert_close(out_flex, out_sdpa, rtol=0, atol=0) + + def test_rejects_doc_ids_shorter_than_the_sequence(self): + """Guards the VLM image-prefix case, where doc_ids would misalign.""" + model = self._model("flex") + tokens = torch.randint(0, 256, (1, 16)) + with pytest.raises(ValueError, match="doc_ids length"), torch.no_grad(): + model(tokens, doc_ids=torch.zeros(1, 8, dtype=torch.long)) + + # --------------------------------------------------------------------------- # MemoryMappedDataset with packing # --------------------------------------------------------------------------- diff --git a/uv.lock b/uv.lock index 0bd1811..9d20e00 100644 --- a/uv.lock +++ b/uv.lock @@ -1444,7 +1444,7 @@ requires-dist = [ { name = "datasets" }, { name = "tensorboard" }, { name = "tokenizers" }, - { name = "torch", specifier = ">=2.4", index = "https://download.pytorch.org/whl/cu128" }, + { name = "torch", specifier = ">=2.6", index = "https://download.pytorch.org/whl/cu128" }, { name = "torchao", specifier = ">=0.17.0" }, { name = "transformers" }, { name = "wandb" },