Skip to content

Add adapter.pre_norm: optional norm ahead of the mlp_2layer projection - #212

Merged
amazloumi merged 4 commits into
mainfrom
feat/adapter-pre-norm
Oct 7, 2026
Merged

amazloumi merged 4 commits into
mainfrom
feat/adapter-pre-norm

Conversation

@amazloumi

@amazloumi amazloumi commented Sep 25, 2026 •

Copy link
Copy Markdown
Member

Summary

  • AdapterConfig.pre_norm: a norm registry key, validated at config time against the registered norms; "" (default) builds nothing, so existing configs keep their state dict and init.
  • MLP2LayerAdapter builds it as ln_q over the vision features and applies it ahead of proj1; reset_parameters re-initializes it through the norm's own reset_parameters(), which RMSNorm now provides. The other adapter types ignore the field.
  • The registry reference lists adapter.pre_norm under the norm category.

Testing

  • uv run ruff check kempnerforge/ tests/ passes
  • uv run ruff format --check kempnerforge/ tests/ scripts/ passes (170 files)
  • uv run pyright kempnerforge/ passes (0 errors)
  • uv run pytest tests/unit/ -v --timeout=60 passes: 1905 passed, 3 skipped (33 new)
  • If distributed code changed: uv run torchrun --nproc_per_node=4 -m pytest tests/distributed/ -v: 102 passed, 2 skipped on each of 4 ranks, same as main
  • If training loop / parallelism / optimizers changed: uv run pytest tests/e2e/ --e2e -v: 27 passed, 4 failed, 1 skipped; the 4 (test_checkpoint_save_and_resume, test_moe_checkpoint_resume, test_pp_checkpoint_save_and_resume, test_sigterm_triggers_emergency_checkpoint) fail on main too
  • uv run pytest tests/integration/ -v --timeout=120: 103 passed, same as main
  • uv run pytest tests/smoke/ --smoke --data-path <pre-tokenized dir> --file-pattern 'shard_*.npy' --data-vocab-size 32000 -v: 19 passed
  • uv run pytest examples/vlm/tests -q: 33 passed; uv run pytest examples/vlm/eval/tests/unit -q: 104 passed
  • uv run sphinx-build -W --keep-going -b html docs docs/_build/html succeeds
  • Diff coverage (--cov --cov-branch): 18/18 changed statements and 9/9 branch arcs in kempnerforge/
  • Defaults unchanged vs main, seed 1234: init state dict, fp32 forward and all gradients byte-identical; a 4-GPU 20-step VLM run is loss-identical at every step
  • pre_norm set, 4-GPU FSDP2: tests/distributed/test_vlm_fsdp.py passes with its adapter set to each registered norm; save → load restores outputs bit-exactly; resuming at step 10 matches the uninterrupted run through step 20
  • Checkpoint mismatch: a no-norm checkpoint loaded into a pre_norm model raises, naming adapter.ln_q.weight; a pre_norm checkpoint loaded into a no-norm model loads and drops the norm, as it does for any parameter-adding field (e.g. model.qk_norm)

Closes #210

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.
adapter.pre_norm is documented by its AdapterConfig field docstring and the changelog line, like the other adapter fields.
@codecov

codecov Bot commented Sep 25, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

Files with missing lines Coverage Δ
kempnerforge/config/adapter.py 100.00% <100.00%> (ø)
kempnerforge/model/adapter.py 97.53% <100.00%> (+0.09%) ⬆️
kempnerforge/model/norm.py 100.00% <100.00%> (ø)
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

…n 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().
@amazloumi
amazloumi requested a review from Naeemkh October 7, 2026 13:37
@amazloumi amazloumi added the core Affecting core label Oct 7, 2026

@Naeemkh Naeemkh left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I verified the compatibility claim rather than taking it on trust: with no norm configured the attribute never enters the module registry, so the state dict and init are unchanged for existing configs. All four builders already swallow unknown keyword arguments, so passing the new field to every adapter type is harmless. Adding the reset method to RMSNorm is also correctly scoped, since the only caller is the adapter re-init on a meta-device build.

One thing worth closing. Config validation checks that the value names a registered norm, but the adapter re-init additionally requires that norm to provide a reset method. Both registered norms do after this PR, so nothing is broken today, but the registry is meant to be extended and the next norm added without one would fail inside FSDP setup instead of at config time. Either guard the call or validate the contract alongside the registry lookup.

Minor: the norm is built with the default epsilon with no way to configure it, while model norms take theirs from config.

@amazloumi
amazloumi merged commit 47b6d67 into main Oct 7, 2026
6 checks passed
@amazloumi
amazloumi deleted the feat/adapter-pre-norm branch October 7, 2026 23:47
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

core Affecting core

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Optional pre-projection norm for the mlp_2layer adapter (adapter.pre_norm)

2 participants