Repository navigation
Add adapter.pre_norm: optional norm ahead of the mlp_2layer projection - #212
Conversation
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 Report✅ All modified and coverable lines are covered by tests.
🚀 New features to boost your workflow:
|
…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().
Naeemkh
left a comment
There was a problem hiding this comment.
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.
Summary
AdapterConfig.pre_norm: anormregistry key, validated at config time against the registered norms;""(default) builds nothing, so existing configs keep their state dict and init.MLP2LayerAdapterbuilds it asln_qover the vision features and applies it ahead ofproj1;reset_parametersre-initializes it through the norm's ownreset_parameters(), whichRMSNormnow provides. The other adapter types ignore the field.adapter.pre_normunder thenormcategory.Testing
uv run ruff check kempnerforge/ tests/passesuv run ruff format --check kempnerforge/ tests/ scripts/passes (170 files)uv run pyright kempnerforge/passes (0 errors)uv run pytest tests/unit/ -v --timeout=60passes: 1905 passed, 3 skipped (33 new)uv run torchrun --nproc_per_node=4 -m pytest tests/distributed/ -v: 102 passed, 2 skipped on each of 4 ranks, same asmainuv 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 onmaintoouv run pytest tests/integration/ -v --timeout=120: 103 passed, same asmainuv run pytest tests/smoke/ --smoke --data-path <pre-tokenized dir> --file-pattern 'shard_*.npy' --data-vocab-size 32000 -v: 19 passeduv run pytest examples/vlm/tests -q: 33 passed;uv run pytest examples/vlm/eval/tests/unit -q: 104 passeduv run sphinx-build -W --keep-going -b html docs docs/_build/htmlsucceeds--cov --cov-branch): 18/18 changed statements and 9/9 branch arcs inkempnerforge/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 steppre_normset, 4-GPU FSDP2:tests/distributed/test_vlm_fsdp.pypasses with its adapter set to each registered norm; save → load restores outputs bit-exactly; resuming at step 10 matches the uninterrupted run through step 20pre_normmodel raises, namingadapter.ln_q.weight; apre_normcheckpoint 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