From 55206f217aad8de0656c86e81d3844670a1ca37a Mon Sep 17 00:00:00 2001 From: amazloumi Date: Fri, 25 Sep 2026 09:40:47 -0400 Subject: [PATCH 1/4] Add adapter.pre_norm: optional norm ahead of the mlp_2layer projection A norm-registry key on AdapterConfig, validated against the registered norms, builds ln_q over the vision features ahead of proj1 so the projection sees unit-scale features; "" builds nothing, leaving the default state dict and projection init unchanged. reset_parameters re-initializes the norm after a meta-device build. --- CHANGELOG.md | 1 + docs/configuration/config-sections.md | 15 +++ kempnerforge/config/adapter.py | 14 +++ kempnerforge/model/adapter.py | 32 ++++-- tests/unit/test_adapter.py | 140 ++++++++++++++++++++++++++ tests/unit/test_vlm.py | 86 ++++++++++++++++ 6 files changed, 282 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 30cc794..19a7db8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added +- **`adapter.pre_norm`**: an optional norm-registry-selected norm (`"rmsnorm"` / `"layernorm"`) over the vision features ahead of `mlp_2layer`'s first projection, exposed as `adapter.ln_q`; `""` (default) builds no module, so existing configs keep their state dict and init unchanged. - **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. diff --git a/docs/configuration/config-sections.md b/docs/configuration/config-sections.md index b0d7a53..00af519 100644 --- a/docs/configuration/config-sections.md +++ b/docs/configuration/config-sections.md @@ -87,6 +87,21 @@ Architecture hyperparameters and MoE knobs. Computed properties: `is_moe`, `head_dim`, `computed_ffn_hidden_dim`, `num_params_estimate`. +## `[adapter]` — `AdapterConfig` + +The vision→LLM connector of a VLM run; read only when `[vlm]` is set, and +omitting it selects `mlp_2layer` with the defaults below. +[`kempnerforge/config/adapter.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/config/adapter.py). + +| Field | Type | Default | Purpose | +|-------|------|---------|---------| +| `type` | `str` | `"mlp_2layer"` | `adapter` registry key: `mlp_2layer` / `linear` keep the token count, `avgpool` / `attentional_pool` pool the patch grid | +| `hidden_dim` | `int` | `0` | `mlp_2layer` hidden width; `0` → `model.dim` | +| `activation` | `str` | `"gelu"` | activation between the `mlp_2layer` projections: `"gelu"`, `"silu"`, `"relu"` | +| `pre_norm` | `str` | `""` | `norm` registry key (`"rmsnorm"` / `"layernorm"`) for a norm over the vision features ahead of `mlp_2layer`'s first projection (`adapter.ln_q`); `""` builds none | +| `pool_window` | `int` | `2` | pooling kernel side for `avgpool` / `attentional_pool` | +| `pool_heads` | `int` | `16` | attention heads for `attentional_pool`; must divide the vision feature dim | + ## `[train]` — `TrainConfig` Training-loop hyperparameters. diff --git a/kempnerforge/config/adapter.py b/kempnerforge/config/adapter.py index 1b0b044..f6041d8 100644 --- a/kempnerforge/config/adapter.py +++ b/kempnerforge/config/adapter.py @@ -29,6 +29,9 @@ class AdapterConfig: ``out_dim``"; ignored by the other types. activation: Activation between the two MLP projections. One of ``"gelu"`` (default), ``"silu"``, ``"relu"``. ``mlp_2layer`` only. + pre_norm: Norm-registry key (``"rmsnorm"`` / ``"layernorm"``) for a norm + over the vision features ahead of ``mlp_2layer``'s first projection + (exposed as ``ln_q``); ``""`` (default) builds none. ``mlp_2layer`` only. pool_window: Pooling kernel side for the pooling adapters (e.g. ``2`` for image 2×2, ``3`` for video 3×3); ignored by projection adapters. pool_heads: Number of attention heads for ``attentional_pool``; must @@ -38,6 +41,7 @@ class AdapterConfig: type: str = "mlp_2layer" hidden_dim: int = 0 activation: str = "gelu" + pre_norm: str = "" pool_window: int = 2 pool_heads: int = 16 @@ -60,6 +64,15 @@ def __post_init__(self) -> None: raise ValueError( f"Unknown adapter.activation: {self.activation!r}. Options: 'gelu', 'silu', 'relu'." ) + if self.pre_norm: + import kempnerforge.model.norm # noqa: F401, PLC0415 (registers the norms) + + norms = tuple(registry.list("norm")) + if self.pre_norm not in norms: + raise ValueError( + f"Unknown adapter.pre_norm: {self.pre_norm!r}. Registered norms: " + f"{sorted(norms)}; '' builds none." + ) if self.pool_window <= 0: raise ValueError(f"adapter.pool_window must be positive (got {self.pool_window})") if self.pool_heads <= 0: @@ -75,6 +88,7 @@ def extra_kwargs(self) -> dict[str, Any]: return { "hidden_dim": self.hidden_dim or None, "activation": self.activation, + "pre_norm": self.pre_norm or None, "pool_window": self.pool_window, "pool_heads": self.pool_heads, } diff --git a/kempnerforge/model/adapter.py b/kempnerforge/model/adapter.py index 7242712..dc148e3 100644 --- a/kempnerforge/model/adapter.py +++ b/kempnerforge/model/adapter.py @@ -7,8 +7,9 @@ Two families: - **Projection adapters** keep the token count (``out_tokens == num_tokens``): - ``mlp_2layer`` (default, the canonical LLaVA-family 2-layer MLP) and - ``linear`` (single ``nn.Linear``, an ablation baseline). + ``mlp_2layer`` (default, the canonical LLaVA-family 2-layer MLP, with an + optional pre-projection norm) and ``linear`` (single ``nn.Linear``, an + ablation baseline). - **Pooling adapters** reduce the token count by pooling the square patch grid before projecting: ``avgpool`` (window-average, the cheapest reducer) and ``attentional_pool`` (Molmo2-style per-window multi-head attention with the @@ -32,6 +33,7 @@ import torch.nn.functional as F from kempnerforge.config.registry import registry +from kempnerforge.model.norm import build_norm _ADAPTER_ACTIVATIONS: dict[str, type[nn.Module]] = { "gelu": nn.GELU, @@ -133,11 +135,13 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: # pragma: no cover class MLP2LayerAdapter(VisionAdapter): """2-layer MLP from image-feature dim to LLM embedding dim. - Architecture: ``Linear(in_dim, hidden) -> activation -> Linear(hidden, out_dim)``. - ``hidden_dim=None`` defaults to ``out_dim``. Keeps the token count. + Architecture: ``[pre_norm ->] Linear(in_dim, hidden) -> activation -> + Linear(hidden, out_dim)``. ``hidden_dim=None`` defaults to ``out_dim``; + ``pre_norm`` (a norm-registry key) adds ``ln_q`` over the vision features + and ``None`` builds none. Keeps the token count. ``reset_parameters`` is provided so callers that materialize adapters - from meta can re-initialize weights with the standard Linear defaults. + from meta can re-initialize weights with the standard defaults. """ def __init__( @@ -146,6 +150,7 @@ def __init__( out_dim: int, hidden_dim: int | None = None, activation: str = "gelu", + pre_norm: str | None = None, ) -> None: super().__init__() if in_dim <= 0 or out_dim <= 0: @@ -155,19 +160,32 @@ def __init__( f"Unknown adapter activation: {activation!r}. Options: {list(_ADAPTER_ACTIVATIONS)}" ) hidden = hidden_dim if hidden_dim and hidden_dim > 0 else out_dim + # Sits ahead of proj1 so the projection sees unit-scale features instead + # of absorbing the encoder's output scale; None builds nothing, leaving + # the default module, state dict and RNG stream untouched. + self.ln_q = build_norm(pre_norm, in_dim) if pre_norm else None self.proj1 = nn.Linear(in_dim, hidden, bias=True) self.act = _ADAPTER_ACTIVATIONS[activation]() self.proj2 = nn.Linear(hidden, out_dim, bias=True) def reset_parameters(self) -> None: - """Re-run ``nn.Linear`` default init on both projections. + """Re-init the projections and the pre-norm, if any. Used after ``to_empty(device=...)`` on a meta-device build. """ + if self.ln_q is not None: + reset = getattr(self.ln_q, "reset_parameters", None) + if callable(reset): + reset() + else: # a bare gain (+ bias) norm such as RMSNorm: weight -> 1, bias -> 0 + for name, param in self.ln_q.named_parameters(): + (nn.init.zeros_ if name.endswith("bias") else nn.init.ones_)(param) self.proj1.reset_parameters() self.proj2.reset_parameters() def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.ln_q is not None: + x = self.ln_q(x) return self.proj2(self.act(self.proj1(x))) @@ -344,6 +362,7 @@ def _build_mlp_2layer( out_dim: int, hidden_dim: int | None = None, activation: str = "gelu", + pre_norm: str | None = None, **_: Any, ) -> VisionAdapter: return MLP2LayerAdapter( @@ -351,6 +370,7 @@ def _build_mlp_2layer( out_dim=out_dim, hidden_dim=hidden_dim, activation=activation, + pre_norm=pre_norm, ) diff --git a/tests/unit/test_adapter.py b/tests/unit/test_adapter.py index 208e8c5..b5c759e 100644 --- a/tests/unit/test_adapter.py +++ b/tests/unit/test_adapter.py @@ -753,3 +753,143 @@ def test_dispatches_to_attentional_with_heads(self): adapter = build_adapter(cfg, in_dim=32, out_dim=16) assert isinstance(adapter, AttentionalPoolAdapter) assert adapter.pool_heads == 8 + + +# --------------------------------------------------------------------------- +# mlp_2layer pre_norm +# --------------------------------------------------------------------------- + +_REGISTERED_NORMS = sorted(registry.list("norm")) +_PROJECTION_KEYS = {"proj1.weight", "proj1.bias", "proj2.weight", "proj2.bias"} + + +def _reference_projections(seed: int, in_dim: int, hidden: int, out_dim: int): + """The two ``nn.Linear`` draws a plain 2-layer MLP makes at ``seed``; every + ``mlp_2layer`` build must reproduce them for its projections.""" + torch.manual_seed(seed) + return torch.nn.Linear(in_dim, hidden), torch.nn.Linear(hidden, out_dim) + + +class TestMLP2LayerPreNorm: + def test_registry_has_norms_to_test(self): + """Refuse rather than pass vacuously if the norm registry is empty.""" + assert _REGISTERED_NORMS + + @pytest.mark.parametrize("kwargs", [{}, {"pre_norm": ""}], ids=["absent", "empty"]) + def test_off_builds_no_module_and_no_state(self, kwargs): + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, **kwargs) + assert adapter.ln_q is None + assert "ln_q" not in dict(adapter.named_modules()) + assert set(adapter.state_dict()) == _PROJECTION_KEYS + + @pytest.mark.parametrize("pre_norm", [None, "", *_REGISTERED_NORMS]) + def test_projections_draw_the_same_rng_as_a_bare_mlp(self, pre_norm): + """Off, the knob builds nothing; on, the registered norms are + constant-initialized -- either way proj1/proj2 get the draw a plain + 2-layer MLP gets, so no existing config's init shifts.""" + ref1, ref2 = _reference_projections(11, 32, 16, 16) + torch.manual_seed(11) + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=pre_norm) + for ours, ref in ((adapter.proj1, ref1), (adapter.proj2, ref2)): + assert torch.equal(ours.weight, ref.weight) + assert torch.equal(ours.bias, ref.bias) + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_norm_is_built_over_the_input_dim(self, norm): + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm) + assert adapter.ln_q is not None + # Every norm parameter spans in_dim: the norm sits on the vision + # features ahead of proj1, not on the hidden width. + assert all(p.shape == (32,) for p in adapter.ln_q.parameters()) + assert set(adapter.state_dict()) == _PROJECTION_KEYS | { + f"ln_q.{k}" for k in adapter.ln_q.state_dict() + } + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_forward_is_norm_then_mlp(self, norm): + torch.manual_seed(0) + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm).to(DEVICE) + x = torch.randn(2, 4, 32, device=DEVICE) * 10.0 # far from unit scale + with torch.no_grad(): + expected = adapter.proj2(adapter.act(adapter.proj1(adapter.ln_q(x)))) + bare = adapter.proj2(adapter.act(adapter.proj1(x))) + out = adapter(x) + assert out.shape == (2, 4, 16) + assert torch.equal(out, expected) + assert not torch.allclose(out, bare) # the norm is on the path, not merely built + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_norm_receives_gradient(self, norm): + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm).to(DEVICE) + adapter(torch.randn(2, 4, 32, device=DEVICE)).sum().backward() + for name, p in adapter.ln_q.named_parameters(): + assert p.grad is not None and torch.isfinite(p.grad).all(), name + assert (adapter.ln_q.weight.grad != 0).any() + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_reset_parameters_restores_the_norm_after_a_meta_build(self, norm): + """The meta-device path runs to_empty() then reset_parameters(); the + norm must come back to its constructor init, not keep whatever memory + to_empty handed it.""" + fresh = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm) + with torch.device("meta"): + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm) + adapter.to_empty(device=torch.device("cpu")) + with torch.no_grad(): + for p in adapter.ln_q.parameters(): + p.fill_(float("nan")) + adapter.reset_parameters() + for k, v in adapter.ln_q.state_dict().items(): + assert torch.equal(v, fresh.ln_q.state_dict()[k]), k + assert all(torch.isfinite(p).all() for p in adapter.parameters()) + + def test_reset_parameters_defers_to_a_norm_that_has_its_own(self): + """A registered norm with its own reset_parameters() is re-initialized + by it, not by the weight->1 / bias->0 fallback.""" + calls: list[str] = [] + + class _Norm(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.full((32,), 0.5)) + + def reset_parameters(self): + calls.append("reset") + + adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm="rmsnorm") + adapter.ln_q = _Norm() + adapter.reset_parameters() + assert calls == ["reset"] + assert torch.equal(adapter.ln_q.weight, torch.full((32,), 0.5)) + + +class TestAdapterConfigPreNorm: + def test_default_is_off(self): + cfg = AdapterConfig() + assert cfg.pre_norm == "" + assert cfg.extra_kwargs()["pre_norm"] is None + assert build_adapter(cfg, in_dim=32, out_dim=16).ln_q is None + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_every_registered_norm_is_accepted_and_built(self, norm): + cfg = AdapterConfig(pre_norm=norm) + assert cfg.extra_kwargs()["pre_norm"] == norm + adapter = build_adapter(cfg, in_dim=32, out_dim=16) + assert isinstance(adapter, MLP2LayerAdapter) + assert adapter.ln_q is not None + assert adapter(torch.randn(1, 4, 32)).shape == (1, 4, 16) + + def test_unknown_norm_is_rejected_naming_the_registered_ones(self): + with pytest.raises(ValueError, match=r"Unknown adapter\.pre_norm: 'nope'") as info: + AdapterConfig(pre_norm="nope") + for norm in _REGISTERED_NORMS: + assert norm in str(info.value) + + @pytest.mark.parametrize("adapter_type", ["linear", "avgpool", "attentional_pool"]) + def test_other_adapter_types_ignore_it(self, adapter_type): + """A shared [adapter] section carrying pre_norm must not break the + builders that have no use for it.""" + cfg = AdapterConfig(type=adapter_type, pre_norm="rmsnorm", pool_heads=8) + adapter = build_adapter(cfg, in_dim=32, out_dim=16) + assert not isinstance(adapter, MLP2LayerAdapter) + assert not hasattr(adapter, "ln_q") diff --git a/tests/unit/test_vlm.py b/tests/unit/test_vlm.py index e9fd04c..1511e34 100644 --- a/tests/unit/test_vlm.py +++ b/tests/unit/test_vlm.py @@ -8,9 +8,13 @@ import pytest import torch +import torch.distributed.checkpoint as dcp +from torch.distributed.checkpoint.api import CheckpointException +from torch.distributed.checkpoint.state_dict import get_model_state_dict, set_model_state_dict from kempnerforge.config.adapter import AdapterConfig from kempnerforge.config.model import ModelConfig +from kempnerforge.config.registry import registry from kempnerforge.config.vision import VisionEncoderConfig from kempnerforge.config.vlm import ( CrossAttentionConfig, @@ -726,3 +730,85 @@ def test_undecodable_clip_stays_finite(self, arch): with torch.no_grad(): logits, _ = w(pix, ids, frame_mask=fm) assert torch.isfinite(logits).all(), f"{arch}: NaN/inf with an all-padded clip" + + +# --------------------------------------------------------------------------- +# adapter.pre_norm through build_vlm_wrapper +# --------------------------------------------------------------------------- + +_REGISTERED_NORMS = sorted(registry.list("norm")) + + +def _pre_norm_wrapper(pre_norm: str = "") -> VLMWrapper: + mc, vc, _, lc = _tiny_configs() + return build_vlm_wrapper(mc, vc, AdapterConfig(pre_norm=pre_norm), lc) + + +class TestVLMAdapterPreNorm: + def test_off_is_the_default_wrapper(self): + torch.manual_seed(0) + default = _build_tiny_wrapper().state_dict() + torch.manual_seed(0) + off = _pre_norm_wrapper("").state_dict() + assert list(off) == list(default) + assert all(torch.equal(off[k], default[k]) for k in default) + assert not [k for k in default if "ln_q" in k] + + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_on_adds_only_the_norm_keys(self, norm): + """Everything else -- the transformer built after the adapter + included -- keeps the same init, so the norm costs no RNG.""" + torch.manual_seed(0) + off = _pre_norm_wrapper("").state_dict() + torch.manual_seed(0) + w = _pre_norm_wrapper(norm) + on = w.state_dict() + assert [k for k in on if k not in off] == [ + f"adapter.ln_q.{k}" for k in w.adapter.ln_q.state_dict() + ] + assert [k for k in off if k not in on] == [] + assert all(torch.equal(on[k], off[k]) for k in off) + + def test_trains_end_to_end(self): + torch.manual_seed(0) + w = _pre_norm_wrapper("rmsnorm").to(DEVICE) + pix = torch.randn(2, 3, 16, 16, device=DEVICE) + ids = torch.randint(0, 256, (2, 16), device=DEVICE) + logits, labels = w(pix, ids, ids.clone()) + assert logits.shape == (2, 16, 256) + assert labels.shape == (2, 16) + logits.float().logsumexp(-1).mean().backward() + grad = w.adapter.ln_q.weight.grad + assert grad is not None and torch.isfinite(grad).all() and (grad != 0).any() + + +class TestVLMAdapterPreNormCheckpoint: + """The same ``get_model_state_dict`` template and ``dcp.load`` call that + ``CheckpointManager.load`` makes, in a single process.""" + + @staticmethod + def _save(wrapper: VLMWrapper, path) -> None: + dcp.save({"model": get_model_state_dict(wrapper)}, checkpoint_id=str(path)) + + def test_round_trip_restores_the_norm(self, tmp_path): + torch.manual_seed(0) + source = _pre_norm_wrapper("layernorm") + with torch.no_grad(): + source.adapter.ln_q.weight.mul_(3.0) + source.adapter.ln_q.bias.add_(0.25) + self._save(source, tmp_path) + torch.manual_seed(1) + target = _pre_norm_wrapper("layernorm") + template = {"model": get_model_state_dict(target)} + dcp.load(template, checkpoint_id=str(tmp_path)) + set_model_state_dict(target, template["model"]) + for k, v in source.state_dict().items(): + assert torch.equal(target.state_dict()[k], v), k + + def test_checkpoint_without_the_norm_refuses_a_model_with_it(self, tmp_path): + """DCP rejects a template key the checkpoint lacks, so a checkpoint + saved with pre_norm off cannot be loaded into a pre_norm model.""" + self._save(_pre_norm_wrapper(""), tmp_path) + target = _pre_norm_wrapper("rmsnorm") + with pytest.raises(CheckpointException, match=r"adapter\.ln_q\.weight"): + dcp.load({"model": get_model_state_dict(target)}, checkpoint_id=str(tmp_path)) From 0da024d8b4a0d8845e1ed64c97d5051a904f8989 Mon Sep 17 00:00:00 2001 From: amazloumi Date: Fri, 25 Sep 2026 09:55:54 -0400 Subject: [PATCH 2/4] Trim adapter.pre_norm comments and changelog line --- CHANGELOG.md | 2 +- kempnerforge/config/adapter.py | 2 +- kempnerforge/model/adapter.py | 11 ++++------- tests/unit/test_adapter.py | 6 ++---- 4 files changed, 8 insertions(+), 13 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 19a7db8..5c53b82 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,7 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added -- **`adapter.pre_norm`**: an optional norm-registry-selected norm (`"rmsnorm"` / `"layernorm"`) over the vision features ahead of `mlp_2layer`'s first projection, exposed as `adapter.ln_q`; `""` (default) builds no module, so existing configs keep their state dict and init unchanged. +- **`adapter.pre_norm`**: an optional norm-registry norm over the vision features ahead of `mlp_2layer`'s first projection (`adapter.ln_q`); `""` (default) builds none, so existing configs are unchanged. - **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. diff --git a/kempnerforge/config/adapter.py b/kempnerforge/config/adapter.py index f6041d8..fdcd1b1 100644 --- a/kempnerforge/config/adapter.py +++ b/kempnerforge/config/adapter.py @@ -65,7 +65,7 @@ def __post_init__(self) -> None: f"Unknown adapter.activation: {self.activation!r}. Options: 'gelu', 'silu', 'relu'." ) if self.pre_norm: - import kempnerforge.model.norm # noqa: F401, PLC0415 (registers the norms) + import kempnerforge.model.norm # noqa: F401, PLC0415 norms = tuple(registry.list("norm")) if self.pre_norm not in norms: diff --git a/kempnerforge/model/adapter.py b/kempnerforge/model/adapter.py index dc148e3..c74d73b 100644 --- a/kempnerforge/model/adapter.py +++ b/kempnerforge/model/adapter.py @@ -7,9 +7,8 @@ Two families: - **Projection adapters** keep the token count (``out_tokens == num_tokens``): - ``mlp_2layer`` (default, the canonical LLaVA-family 2-layer MLP, with an - optional pre-projection norm) and ``linear`` (single ``nn.Linear``, an - ablation baseline). + ``mlp_2layer`` (default, the canonical LLaVA-family 2-layer MLP) and + ``linear`` (single ``nn.Linear``, an ablation baseline). - **Pooling adapters** reduce the token count by pooling the square patch grid before projecting: ``avgpool`` (window-average, the cheapest reducer) and ``attentional_pool`` (Molmo2-style per-window multi-head attention with the @@ -160,9 +159,7 @@ def __init__( f"Unknown adapter activation: {activation!r}. Options: {list(_ADAPTER_ACTIVATIONS)}" ) hidden = hidden_dim if hidden_dim and hidden_dim > 0 else out_dim - # Sits ahead of proj1 so the projection sees unit-scale features instead - # of absorbing the encoder's output scale; None builds nothing, leaving - # the default module, state dict and RNG stream untouched. + # Ahead of proj1, so the projection sees normalized features whatever the encoder's scale. self.ln_q = build_norm(pre_norm, in_dim) if pre_norm else None self.proj1 = nn.Linear(in_dim, hidden, bias=True) self.act = _ADAPTER_ACTIVATIONS[activation]() @@ -177,7 +174,7 @@ def reset_parameters(self) -> None: reset = getattr(self.ln_q, "reset_parameters", None) if callable(reset): reset() - else: # a bare gain (+ bias) norm such as RMSNorm: weight -> 1, bias -> 0 + else: # RMSNorm has no reset_parameters() for name, param in self.ln_q.named_parameters(): (nn.init.zeros_ if name.endswith("bias") else nn.init.ones_)(param) self.proj1.reset_parameters() diff --git a/tests/unit/test_adapter.py b/tests/unit/test_adapter.py index b5c759e..dc83165 100644 --- a/tests/unit/test_adapter.py +++ b/tests/unit/test_adapter.py @@ -798,8 +798,6 @@ def test_projections_draw_the_same_rng_as_a_bare_mlp(self, pre_norm): def test_norm_is_built_over_the_input_dim(self, norm): adapter = MLP2LayerAdapter(in_dim=32, out_dim=16, pre_norm=norm) assert adapter.ln_q is not None - # Every norm parameter spans in_dim: the norm sits on the vision - # features ahead of proj1, not on the hidden width. assert all(p.shape == (32,) for p in adapter.ln_q.parameters()) assert set(adapter.state_dict()) == _PROJECTION_KEYS | { f"ln_q.{k}" for k in adapter.ln_q.state_dict() @@ -812,11 +810,11 @@ def test_forward_is_norm_then_mlp(self, norm): x = torch.randn(2, 4, 32, device=DEVICE) * 10.0 # far from unit scale with torch.no_grad(): expected = adapter.proj2(adapter.act(adapter.proj1(adapter.ln_q(x)))) - bare = adapter.proj2(adapter.act(adapter.proj1(x))) + without_norm = adapter.proj2(adapter.act(adapter.proj1(x))) out = adapter(x) assert out.shape == (2, 4, 16) assert torch.equal(out, expected) - assert not torch.allclose(out, bare) # the norm is on the path, not merely built + assert not torch.allclose(out, without_norm) @pytest.mark.parametrize("norm", _REGISTERED_NORMS) def test_norm_receives_gradient(self, norm): From 772eac3a79b7e70cc069ef0f1651ea4bd6821261 Mon Sep 17 00:00:00 2001 From: amazloumi Date: Fri, 25 Sep 2026 10:30:38 -0400 Subject: [PATCH 3/4] Drop the [adapter] section from the config reference adapter.pre_norm is documented by its AdapterConfig field docstring and the changelog line, like the other adapter fields. --- docs/configuration/config-sections.md | 15 --------------- 1 file changed, 15 deletions(-) diff --git a/docs/configuration/config-sections.md b/docs/configuration/config-sections.md index 00af519..b0d7a53 100644 --- a/docs/configuration/config-sections.md +++ b/docs/configuration/config-sections.md @@ -87,21 +87,6 @@ Architecture hyperparameters and MoE knobs. Computed properties: `is_moe`, `head_dim`, `computed_ffn_hidden_dim`, `num_params_estimate`. -## `[adapter]` — `AdapterConfig` - -The vision→LLM connector of a VLM run; read only when `[vlm]` is set, and -omitting it selects `mlp_2layer` with the defaults below. -[`kempnerforge/config/adapter.py`](https://github.com/KempnerInstitute/KempnerForge/blob/main/kempnerforge/config/adapter.py). - -| Field | Type | Default | Purpose | -|-------|------|---------|---------| -| `type` | `str` | `"mlp_2layer"` | `adapter` registry key: `mlp_2layer` / `linear` keep the token count, `avgpool` / `attentional_pool` pool the patch grid | -| `hidden_dim` | `int` | `0` | `mlp_2layer` hidden width; `0` → `model.dim` | -| `activation` | `str` | `"gelu"` | activation between the `mlp_2layer` projections: `"gelu"`, `"silu"`, `"relu"` | -| `pre_norm` | `str` | `""` | `norm` registry key (`"rmsnorm"` / `"layernorm"`) for a norm over the vision features ahead of `mlp_2layer`'s first projection (`adapter.ln_q`); `""` builds none | -| `pool_window` | `int` | `2` | pooling kernel side for `avgpool` / `attentional_pool` | -| `pool_heads` | `int` | `16` | attention heads for `attentional_pool`; must divide the vision feature dim | - ## `[train]` — `TrainConfig` Training-loop hyperparameters. From 82d30058199f2e91eebd5587fdcedb0c7bb50710 Mon Sep 17 00:00:00 2001 From: amazloumi Date: Fri, 25 Sep 2026 12:17:25 -0400 Subject: [PATCH 4/4] Re-init the adapter pre-norm through its own reset_parameters; tighten the pre_norm tests RMSNorm gains reset_parameters(), so MLP2LayerAdapter re-initializes any registered norm by calling it instead of guessing from parameter names. Tests: other adapter types are unchanged by pre_norm (state and output, ragged grid included), a norm registered at runtime is accepted, the unknown-name check covers a non-mlp type, and every registered norm exposes reset_parameters(). --- CHANGELOG.md | 2 +- docs/configuration/registry.md | 2 +- kempnerforge/model/adapter.py | 8 +---- kempnerforge/model/norm.py | 3 ++ tests/unit/test_adapter.py | 53 +++++++++++++++++++++++++++------- tests/unit/test_vlm.py | 15 +++++----- 6 files changed, 55 insertions(+), 28 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5c53b82..30b31b3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,7 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added -- **`adapter.pre_norm`**: an optional norm-registry norm over the vision features ahead of `mlp_2layer`'s first projection (`adapter.ln_q`); `""` (default) builds none, so existing configs are unchanged. +- **`adapter.pre_norm`**: an optional norm-registry norm over the vision features ahead of `mlp_2layer`'s first projection (`adapter.ln_q`), re-initialized through the norm's `reset_parameters()` (added to `RMSNorm`); `""` (default) builds none, so existing configs are unchanged. - **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. diff --git a/docs/configuration/registry.md b/docs/configuration/registry.md index 5a60cdc..3439808 100644 --- a/docs/configuration/registry.md +++ b/docs/configuration/registry.md @@ -27,7 +27,7 @@ raises `KeyError` with an "Available: […]" hint when a name is missing; | `optimizer` | `adamw`, `lion`, `muon`, `schedule_free_adamw` | `optimizer.name` | | `scheduler` | `cosine`, `linear`, `wsd`, `constant`, `rex`, `none` | `scheduler.name` | | `loss` | `cross_entropy`, `chunked_cross_entropy` | `train.loss_fn` | -| `norm` | `rmsnorm`, `layernorm` | `model.norm_type` | +| `norm` | `rmsnorm`, `layernorm` | `model.norm_type`, `adapter.pre_norm` | | `router` | `softmax_topk`, `sigmoid_topk` | `model.moe_router` | | `mlp` | `swiglu`, `standard_gelu`, `standard_relu` | `model.activation` (mapped) | diff --git a/kempnerforge/model/adapter.py b/kempnerforge/model/adapter.py index c74d73b..2da8e0f 100644 --- a/kempnerforge/model/adapter.py +++ b/kempnerforge/model/adapter.py @@ -159,7 +159,6 @@ def __init__( f"Unknown adapter activation: {activation!r}. Options: {list(_ADAPTER_ACTIVATIONS)}" ) hidden = hidden_dim if hidden_dim and hidden_dim > 0 else out_dim - # Ahead of proj1, so the projection sees normalized features whatever the encoder's scale. self.ln_q = build_norm(pre_norm, in_dim) if pre_norm else None self.proj1 = nn.Linear(in_dim, hidden, bias=True) self.act = _ADAPTER_ACTIVATIONS[activation]() @@ -171,12 +170,7 @@ def reset_parameters(self) -> None: Used after ``to_empty(device=...)`` on a meta-device build. """ if self.ln_q is not None: - reset = getattr(self.ln_q, "reset_parameters", None) - if callable(reset): - reset() - else: # RMSNorm has no reset_parameters() - for name, param in self.ln_q.named_parameters(): - (nn.init.zeros_ if name.endswith("bias") else nn.init.ones_)(param) + self.ln_q.reset_parameters() # type: ignore[reportCallIssue] self.proj1.reset_parameters() self.proj2.reset_parameters() diff --git a/kempnerforge/model/norm.py b/kempnerforge/model/norm.py index 59f2c2c..00e6d0f 100644 --- a/kempnerforge/model/norm.py +++ b/kempnerforge/model/norm.py @@ -19,6 +19,9 @@ def __init__(self, dim: int, eps: float = 1e-5) -> None: self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps + def reset_parameters(self) -> None: + nn.init.ones_(self.weight) + def forward(self, x: torch.Tensor) -> torch.Tensor: # float32 for numerical stability, then cast back dtype = x.dtype diff --git a/tests/unit/test_adapter.py b/tests/unit/test_adapter.py index dc83165..68e76d5 100644 --- a/tests/unit/test_adapter.py +++ b/tests/unit/test_adapter.py @@ -23,6 +23,7 @@ build_adapter, pooled_token_count, ) +from kempnerforge.model.norm import build_norm DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -824,6 +825,11 @@ def test_norm_receives_gradient(self, norm): assert p.grad is not None and torch.isfinite(p.grad).all(), name assert (adapter.ln_q.weight.grad != 0).any() + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) + def test_every_registered_norm_has_reset_parameters(self, norm): + """The adapter re-initializes its norm by calling reset_parameters().""" + assert callable(getattr(build_norm(norm, 32), "reset_parameters", None)) + @pytest.mark.parametrize("norm", _REGISTERED_NORMS) def test_reset_parameters_restores_the_norm_after_a_meta_build(self, norm): """The meta-device path runs to_empty() then reset_parameters(); the @@ -841,9 +847,8 @@ def test_reset_parameters_restores_the_norm_after_a_meta_build(self, norm): assert torch.equal(v, fresh.ln_q.state_dict()[k]), k assert all(torch.isfinite(p).all() for p in adapter.parameters()) - def test_reset_parameters_defers_to_a_norm_that_has_its_own(self): - """A registered norm with its own reset_parameters() is re-initialized - by it, not by the weight->1 / bias->0 fallback.""" + def test_reset_parameters_delegates_to_the_norm(self): + """The norm's own reset_parameters() is the only init applied to it.""" calls: list[str] = [] class _Norm(torch.nn.Module): @@ -861,6 +866,18 @@ def reset_parameters(self): assert torch.equal(adapter.ln_q.weight, torch.full((32,), 0.5)) +class _PluginNorm(torch.nn.LayerNorm): + pass + + +@pytest.fixture +def plugin_norm(): + name = "adapter_pre_norm_test_plugin" + registry.register("norm", name, _PluginNorm) + yield name + registry._get_store("norm").pop(name) + + class TestAdapterConfigPreNorm: def test_default_is_off(self): cfg = AdapterConfig() @@ -877,17 +894,31 @@ def test_every_registered_norm_is_accepted_and_built(self, norm): assert adapter.ln_q is not None assert adapter(torch.randn(1, 4, 32)).shape == (1, 4, 16) - def test_unknown_norm_is_rejected_naming_the_registered_ones(self): + def test_a_norm_registered_at_runtime_is_accepted_and_built(self, plugin_norm): + adapter = build_adapter(AdapterConfig(pre_norm=plugin_norm), in_dim=32, out_dim=16) + assert type(adapter.ln_q) is _PluginNorm + + @pytest.mark.parametrize("adapter_type", ["mlp_2layer", "avgpool"]) + def test_unknown_norm_is_rejected_naming_the_registered_ones(self, adapter_type): with pytest.raises(ValueError, match=r"Unknown adapter\.pre_norm: 'nope'") as info: - AdapterConfig(pre_norm="nope") + AdapterConfig(type=adapter_type, pre_norm="nope") for norm in _REGISTERED_NORMS: assert norm in str(info.value) @pytest.mark.parametrize("adapter_type", ["linear", "avgpool", "attentional_pool"]) def test_other_adapter_types_ignore_it(self, adapter_type): - """A shared [adapter] section carrying pre_norm must not break the - builders that have no use for it.""" - cfg = AdapterConfig(type=adapter_type, pre_norm="rmsnorm", pool_heads=8) - adapter = build_adapter(cfg, in_dim=32, out_dim=16) - assert not isinstance(adapter, MLP2LayerAdapter) - assert not hasattr(adapter, "ln_q") + """pre_norm changes neither the state nor the output of the adapters + that have no use for it; 25 tokens also takes the ragged pooling path.""" + x = torch.randn(2, 25, 32) * 10.0 + built = [] + for pre_norm in ("", "rmsnorm"): + torch.manual_seed(0) + cfg = AdapterConfig(type=adapter_type, pre_norm=pre_norm, pool_heads=8) + built.append(build_adapter(cfg, in_dim=32, out_dim=16)) + off, on = built + assert not isinstance(on, MLP2LayerAdapter) + if adapter_type != "linear": + assert _pad_grid_to_windows(x, off.pool_window)[2] is not None + assert list(on.state_dict()) == list(off.state_dict()) + with torch.no_grad(): + assert torch.equal(on(x), off(x)) diff --git a/tests/unit/test_vlm.py b/tests/unit/test_vlm.py index 1511e34..a313cea 100644 --- a/tests/unit/test_vlm.py +++ b/tests/unit/test_vlm.py @@ -745,14 +745,13 @@ def _pre_norm_wrapper(pre_norm: str = "") -> VLMWrapper: class TestVLMAdapterPreNorm: - def test_off_is_the_default_wrapper(self): - torch.manual_seed(0) - default = _build_tiny_wrapper().state_dict() - torch.manual_seed(0) - off = _pre_norm_wrapper("").state_dict() - assert list(off) == list(default) - assert all(torch.equal(off[k], default[k]) for k in default) - assert not [k for k in default if "ln_q" in k] + def test_off_adds_no_adapter_state(self): + assert [k for k in _pre_norm_wrapper("").state_dict() if k.startswith("adapter.")] == [ + "adapter.proj1.weight", + "adapter.proj1.bias", + "adapter.proj2.weight", + "adapter.proj2.bias", + ] @pytest.mark.parametrize("norm", _REGISTERED_NORMS) def test_on_adds_only_the_norm_keys(self, norm):