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
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 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.
Expand Down
2 changes: 1 addition & 1 deletion docs/configuration/registry.md
Original file line number Diff line number Diff line change
Expand Up @@ -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) |

Expand Down
14 changes: 14 additions & 0 deletions kempnerforge/config/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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

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:
Expand All @@ -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,
}
Expand Down
19 changes: 15 additions & 4 deletions kempnerforge/model/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,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,
Expand Down Expand Up @@ -133,11 +134,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__(
Expand All @@ -146,6 +149,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:
Expand All @@ -155,19 +159,24 @@ 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
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:
self.ln_q.reset_parameters() # type: ignore[reportCallIssue]
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)))


Expand Down Expand Up @@ -344,13 +353,15 @@ def _build_mlp_2layer(
out_dim: int,
hidden_dim: int | None = None,
activation: str = "gelu",
pre_norm: str | None = None,
**_: Any,
) -> VisionAdapter:
return MLP2LayerAdapter(
in_dim=in_dim,
out_dim=out_dim,
hidden_dim=hidden_dim,
activation=activation,
pre_norm=pre_norm,
)


Expand Down
3 changes: 3 additions & 0 deletions kempnerforge/model/norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
169 changes: 169 additions & 0 deletions tests/unit/test_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down Expand Up @@ -753,3 +754,171 @@ 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
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))))
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, without_norm)

@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_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
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_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):
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 _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()
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_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(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):
"""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))
Loading
Loading